skip to content

What do you give up by subclassing keras.Model instead of using the Functional API?

level: seniorimportance: should knowfreq 54%

answer

  1. a graph is data, code is not
  2. nothing to inspect, nothing to slice
  3. shape errors move to runtime
  4. serialization stops being automatic
  5. subclass the layer, not the model

basics

~20 s

You 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 s

Functional 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 lines
python
import 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

for a junior

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.

for a middle

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.

for a senior

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.

for a principal

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

context