What do you give up by subclassing keras.Model instead of using the Functional API?
answer
- a graph is data, code is not
- nothing to inspect, nothing to slice
- shape errors move to runtime
- serialization stops being automatic
- subclass the layer, not the model
basics
~20 sYou give up the static layer graph. Without it there is no shape inference before the first batch, summary and plotting are degraded, get_layer(name).output surgery is impossible, and the model no longer serializes from configuration alone. In return you get arbitrary Python control flow in the forward pass.
solid answer
~50 sFunctional models are data: a recorded graph of layer calls that Keras can inspect, slice, draw and reconstruct. A subclassed model is code, and Keras cannot inspect code. Concretely you lose up-front shape propagation (mistakes surface on the first forward pass instead of at definition time), full `model.summary()` output shapes, meaningful `keras.utils.plot_model`, and the ability to build a feature extractor or splice on a new head via `model.get_layer(...).output`. Serialization stops being automatic too: a subclassed model needs `get_config`/`from_config` and registration to round-trip, whereas a Functional topology saves itself. What you buy is expressiveness — data-dependent branching, loops and recursion in `call()` that a static graph cannot represent. In Keras 3 there is a third cost worth naming: a `call()` written against one backend's raw ops is no longer portable, so backend-agnostic code has to go through `keras.ops`. The usual resolution is not to choose one API but to subclass `keras.layers.Layer` for the novel block and assemble the model itself functionally.
code
python · 22 linesimport keras
class GatedFusion(keras.layers.Layer):
def __init__(self, units, **kwargs):
super().__init__(**kwargs)
self.gate = keras.layers.Dense(units, activation="sigmoid")
self.proj = keras.layers.Dense(units)
def call(self, inputs):
a, b = inputs
g = self.gate(keras.layers.Concatenate()([a, b]))
return g * self.proj(a) + (1.0 - g) * self.proj(b)
left = keras.Input(shape=(16,), name="left")
right = keras.Input(shape=(16,), name="right")
fused = GatedFusion(32, name="fusion")([left, right])
out = keras.layers.Dense(1, name="score")(fused)
model = keras.Model(inputs=[left, right], outputs=out)
model.summary()go deeper
Know the headline: subclassing means the architecture lives in Python code, so Keras cannot show you shapes or a diagram until the model runs, and you gain the freedom to write any forward pass you like.
Derive the losses from one fact — a Functional model is an inspectable graph and a subclassed one is opaque code. Name shape inference, summary output shapes, get_layer surgery and automatic serialization as the concrete casualties.
Show judgment: pick the API from whether the forward pass depends on values or only on shapes, and offer the hybrid — custom Layer inside a Functional model — as the pattern that keeps expressiveness without losing tooling.
Own the portability and maintenance angle. Imperative models shift enforcement from the framework to the team, and in a multi-backend Keras 3 world an unconstrained call() can quietly pin a whole codebase to one backend.
## The core distinction: a model as data versus a model as code A Functional model is a **declarative artifact**. Its topology is a recorded graph — which layer was applied to which tensor — and that graph is inspectable, serializable and sliceable. A subclassed model is an **imperative program**. The architecture exists only as the sequence of statements Python executes when `call()` runs. Every capability in the list below follows from this one distinction, and being able to derive them from it, rather than reciting them, is what separates a strong answer from a memorised one. ## What you actually lose **Shape inference before execution.** Functional models propagate shapes from `keras.Input` at construction time, so a wrong hidden dimension is a construction-time error. Subclassed models learn their shapes by running, so the same mistake costs you a forward pass — and in a distributed or long-warm-up setting, considerably more than that. **Full introspection.** `model.summary()` on a subclassed model reports layers and parameter counts once built, but frequently cannot report per-layer output shapes, because no node graph links one layer's output to the next layer's input. `keras.utils.plot_model` has correspondingly little to draw. For a model that other people review, this is a real loss of communicable structure. **Graph surgery.** `model.get_layer("block3").output` returns nothing on a subclassed model — the layer was only called inside Python. Feature extraction, perceptual losses, cutting a pretrained backbone at a chosen depth and attaching a new head: all of these depend on the recorded graph. **Free serialization.** A Functional topology is configuration, so it round-trips without help. A subclassed model's architecture is code, so restoring it requires that code to be importable, plus `get_config`/`from_config` and serialization registration if it is to be reconstructed from a saved file rather than rebuilt by hand. **Structural validation.** The Functional API rejects impossible wiring as you write it — mismatched merges, disconnected inputs. Subclassing accepts anything Python accepts and defers judgement to runtime. ## What you gain **Data-dependent control flow.** A branch taken on the basis of a runtime value, a loop whose count depends on the input, recursion, an early exit — none of these can be expressed as a static graph of layer calls, and all of them are ordinary Python in `call()`. **Multiple, stateful forward paths.** Models whose forward pass differs structurally between phases — some multi-stage generative and reinforcement-learning setups, models with internal memory or per-step caches — are far more natural to write imperatively. **Readability for genuinely algorithmic models.** When the model *is* an algorithm rather than a topology, code that reads like the algorithm beats a graph assembled out of merge layers. ## The Keras 3 portability angle Keras 3 runs on TensorFlow, JAX and PyTorch, chosen by the `KERAS_BACKEND` environment variable. A Functional model is portable by construction, because it is only ever built out of Keras layers. A subclassed `call()` is portable only if it stays within Keras's own operations — `keras.ops` is the backend-agnostic numerics layer for exactly this. The moment `call()` reaches for a raw backend tensor API, the model is pinned to that backend. This is a cost that did not exist in Keras 2 and it is worth naming explicitly, because it turns a stylistic choice into a portability decision. ## The pattern that resolves the tension Almost no real model needs a subclassed *Model*. What it usually needs is one novel *block* — a gated fusion, an unusual attention variant, a custom normalisation. Subclass `keras.layers.Layer` for that block, then wire it into a Functional graph like any other layer. You keep the arbitrary Python where you need it, and you keep shape inference, summary, plotting, surgery and serialisation for the model as a whole. This hybrid is the answer an experienced practitioner gives, and offering it unprompted is a strong signal. ## How to decide Ask what the forward pass depends on. If it depends only on the *shape* of the input, the graph is static and Functional is strictly better. If it depends on the *values* — a decision made per batch — you need the imperative path, and you should then pay the cost deliberately: build the model explicitly rather than lazily, write the serialization hooks, and document the expected input shapes that the framework no longer enforces for you.
- When is subclassing keras.Model genuinely the right call?When the forward pass depends on runtime *values* rather than only on shapes — a branch chosen per batch, a loop whose length varies, recursion, or a staged pass whose structure differs between phases. A static graph of layer calls cannot represent those. If the architecture is fixed once you know the input shape, the Functional API is strictly the better tool.
- In Keras 3, what makes a subclassed call() non-portable across backends?Reaching for a raw backend tensor API inside `call()` pins the model to that backend. Keras 3 supports TensorFlow, JAX and PyTorch behind the same layer API, so backend-agnostic custom code has to be written against `keras.ops`, Keras's own numerics layer. A Functional model built only from Keras layers is portable without any care at all.
- How would you recover model.summary() output shapes for an existing subclassed model?You largely cannot without restructuring: the shapes are missing because there is no recorded node graph, not because of a setting. Build the model so parameter counts appear, and if the per-layer topology genuinely matters, move the custom logic into a Layer subclass and reassemble the model with the Functional API, which restores full shape reporting.
saying these in an interview costs you the question
- Saying subclassed models are inherently slower to train
- Claiming subclassed models cannot use fit() at all
- Believing a subclassed model saves and reloads with no extra work
- Treating subclassing as the professional default and Functional as beginner-only
- Assuming Keras 3 custom call() code is portable regardless of which ops it uses