skip to content

In a custom Keras 3 Layer, what goes in build() versus call()?

level: middleimportance: must knowfreq 75%

answer

  1. three hooks, three different jobs
  2. weights need a shape first
  3. deferred until the first input arrives
  4. add_weight lives in build(input_shape)
  5. compiled backends need variables pre-existing

basics

~10 s

build(input_shape) creates the weights, and Keras runs it once on the first input, when the shape is finally known. call(inputs) does only the forward computation. Creating weights inside call() is a bug.

solid answer

~40 s

A `keras.layers.Layer` has three hooks. `__init__` stores configuration (units, rate, activation) and may build sub-layers, but nothing that depends on the input shape. `build(input_shape)` is where you call `self.add_weight(...)`, because that is the first moment Keras knows the incoming feature dimension — typically `input_shape[-1]`. `call(inputs)` reads those weights and returns the output, using `keras.ops` so it works on any backend. Keras invokes `build` exactly once, on the first `__call__` (or during symbolic wiring from `keras.Input`), and marks the layer built for you. That deferred build is why you can write `Dense(64)` without stating the input size. Creating variables inside `call()` breaks it: they would be recreated on every call and never trained, and compiled backends like JAX require every variable to exist before the traced step function runs.

code

python · 30 lines
python
import keras


class Linear(keras.layers.Layer):
    def __init__(self, units, **kwargs):
        super().__init__(**kwargs)
        self.units = units

    def build(self, input_shape):
        self.kernel = self.add_weight(
            shape=(input_shape[-1], self.units),
            initializer="glorot_uniform",
            trainable=True,
            name="kernel",
        )
        self.bias = self.add_weight(
            shape=(self.units,),
            initializer="zeros",
            trainable=True,
            name="bias",
        )

    def call(self, inputs):
        return keras.ops.matmul(inputs, self.kernel) + self.bias


layer = Linear(4)
print(layer.weights)                      # [] - not built yet
layer(keras.ops.zeros((2, 8)))
print([w.shape for w in layer.weights])   # [(8, 4), (4,)]

go deeper

for a junior

Know the names and the order: __init__ for settings, build for weights, call for the forward pass. Be able to say why Dense(64) needs no input size.

for a middle

Explain the deferred-build lifecycle end to end — first call, shape inspection, one build, then call — and use input_shape[-1] correctly in a add_weight shape you write on the spot.

for a senior

Show why weight creation inside call fails under a compiled backend, not just why it is untidy, and demonstrate validating shapes and raising a useful error in build.

for a principal

Be ready to argue where a team's reusable-layer boundary sits: what belongs in a custom layer versus composed built-ins, and how deferred build interacts with export, shape contracts and review burden across a shared model library.

## The three methods of a Layer Subclassing `keras.layers.Layer` normally means writing three methods, and each has a distinct job: - `__init__(self, ...)` — store hyperparameters (`self.units = units`), instantiate sub-layers (`self.dense = keras.layers.Dense(units)`), call `super().__init__(**kwargs)` so name/dtype handling works. Nothing here may depend on the shape of the data, because no data has arrived. - `build(self, input_shape)` — create the layer's own variables with `self.add_weight(...)`. `input_shape` is a tuple such as `(None, 10, 8)`, where `None` is the batch axis. The feature dimension you need is almost always `input_shape[-1]`. - `call(self, inputs)` — the forward pass. Read `self.kernel`, do the math, return a tensor. ## Why build() exists at all Keras layers are declared by their *output* size, not their input size: `keras.layers.Dense(64)` says "produce 64 features" and says nothing about what arrives. That is only possible if weight creation is deferred until a real (or symbolic) input shows up. The lifecycle is: first `__call__` → Keras inspects the input shape → calls `build(input_shape)` once → sets `self.built = True` → runs `call`. You do not need to set `built` yourself; Keras 3 does it around your `build`. The payoff is that you can stack layers without hand-computing intermediate dimensions, and the same layer class works on `(batch, features)` and `(batch, time, features)` inputs. The cost is that shape errors surface at the first call rather than at construction, which is why `model.summary()` on an unbuilt subclassed model shows unknown shapes until you call it or `model.build(input_shape)` explicitly. ## What add_weight() buys you `self.add_weight(shape=..., initializer=..., dtype=..., trainable=..., regularizer=..., constraint=..., name=...)` returns a `keras.Variable` and, crucially, *registers* it with the layer. A registered variable appears in `layer.weights` (split into `layer.trainable_weights` and `layer.non_trainable_weights`), is handed to the optimizer when trainable, is saved and restored with the model, and is moved/managed by the backend for you. A plain Python attribute holding a raw backend tensor gets none of that. Initializers are given by string (`"glorot_uniform"`, `"zeros"`) or object (`keras.initializers.RandomNormal(stddev=0.02)`). ## Why creating weights in call() is a real bug, not a style nit Two independent failures: 1. **Semantics.** `call` runs on every batch. Variables created there are fresh each time — the layer would never learn anything, and memory grows. 2. **Compilation.** Keras 3 runs on TensorFlow, JAX or PyTorch. The JAX backend executes your model as a pure function over an explicit list of variables; the variables must already exist before the function is traced. Creating one mid-trace is not something the backend can express. TensorFlow's `tf.function` has the same objection to creating variables on a non-first trace. The same reasoning explains why sub-layers created inside `call()` are a mistake: they would be new, untrained sub-layers on every batch and would never be tracked as part of the parent's weights. ## Shape-dependent weights and validation `build` is also the natural place to validate: raise a clear `ValueError` if `input_shape[-1] is None` or if the rank is wrong. Once built, a layer is fixed — feeding it a different last dimension later raises an error rather than silently rebuilding, which is exactly the behaviour you want. ## compute_output_shape For most layers Keras can infer output shapes by running `call` symbolically on `KerasTensor`s, so you write nothing. If your layer does something the symbolic pass cannot follow, implement `compute_output_shape(self, input_shape)` returning the output tuple so the Functional API can wire the graph without executing it. ## Backend-agnostic math Inside `call`, use `keras.ops` (`keras.ops.matmul`, `keras.ops.mean`, `keras.ops.softmax`) rather than `tf.*` or `torch.*`. Keras 3 is multi-backend; a layer written against one backend's ops is a Keras 2-era layer and will not run under `KERAS_BACKEND=jax`. ## Checklist for a custom layer Configuration in `__init__`; `add_weight` in `build(input_shape)`; math in `call` via `keras.ops`; no variable or sub-layer creation in `call`; `input_shape[-1]` for the feature dim; let Keras set `built`.

  • How would you force a subclassed layer to be built before you ever feed it data?
    Call `layer.build(input_shape)` directly, or push a symbolic tensor through it — `layer(keras.Input(shape=(8,)))` — or call `model.build(input_shape=(None, 8))` on the enclosing model. Any of these runs `build` with the shape you name, creates the variables, and makes `model.summary()` show real parameter counts instead of unknowns.
  • Where do sub-layers get created, and when are their weights made?
    Sub-layers are instantiated in `__init__` and assigned to attributes, which is how Keras tracks them. Their own weights are still created lazily: when the parent's `call` first passes a tensor into a sub-layer, that sub-layer's `build` runs. So the parent's `weights` list is empty until the parent has been called once.
  • Why should call() use keras.ops instead of the backend's own ops?
    Keras 3 runs on TensorFlow, JAX or PyTorch, chosen by the `KERAS_BACKEND` environment variable. `keras.ops` is the backend-agnostic numerics layer, so a layer written with it runs unchanged everywhere. Reaching for `tf.*` inside `call` silently pins the layer to one backend and breaks the moment someone switches.

saying these in an interview costs you the question

  • Creating weights in __init__ with a hardcoded input dimension
  • Calling add_weight inside call() on every batch
  • Thinking build() runs once per batch rather than once per layer
  • Setting self.built = True by hand and skipping weight creation
  • Using tf.matmul inside call() in a Keras 3 layer

context