In a custom Keras Layer, how do you hold state that gradients never update?
answer
- state that is not a parameter
- two buckets in layer.weights
- a Python float is not tracked
- add_weight has a trainable flag
- assign and assign_add, not +=
basics
~20 sCreate it with self.add_weight(..., trainable=False) and mutate it with .assign() or .assign_add(). It becomes a tracked keras.Variable listed under layer.non_trainable_weights, saved with the model, and ignored by the optimizer — exactly how BatchNormalization keeps its moving statistics.
solid answer
~40 sUse `self.add_weight(shape=..., initializer="zeros", trainable=False, name=...)` inside `build()`. That returns a `keras.Variable` registered with the layer, so it shows up in `layer.non_trainable_weights` and `layer.weights`, travels with the model when it is saved and restored, and is skipped by the optimizer. You update it inside `call()` with `self.var.assign(value)` or `self.var.assign_add(delta)`. The reason not to use a plain Python attribute — `self.total = 0.0` — is that Keras tracks nothing about it: it is not part of the layer's state, it is not saved, and mutating it inside `call` does not survive a compiled or traced execution, where `call` runs symbolically rather than once per batch. The canonical built-in example is `BatchNormalization`, whose `moving_mean` and `moving_variance` are exactly this: real variables, accumulated during training, never touched by gradient descent.
code
python · 23 linesimport keras
class BatchCounter(keras.layers.Layer):
def build(self, input_shape):
self.seen = self.add_weight(
shape=(),
initializer="zeros",
trainable=False,
name="seen",
)
def call(self, inputs, training=None):
if training:
self.seen.assign_add(1.0)
return inputs
layer = BatchCounter()
layer(keras.ops.zeros((2, 3)), training=True)
print(layer.trainable_weights) # []
print(len(layer.non_trainable_weights)) # 1
print(float(layer.seen)) # 1.0go deeper
Know that a layer can hold values that are not learned by the optimizer, and that add_weight takes a trainable argument that decides which bucket a variable lands in.
Write the add_weight(trainable=False) call from memory, update it with assign_add, and explain why a Python attribute fails under a traced or compiled backend.
Reason about the whole lifecycle: what gets checkpointed, how freezing changes the reported weight split, and how a running statistic must be gated on training mode to stay correct at serving.
Decide the state policy for a shared layer library — what counts as model state versus per-run telemetry, what must round-trip through a checkpoint, and how that state behaves under distribution and export.
## Two kinds of layer state A Keras layer's variables split into two buckets: - **Trainable weights** — created with `add_weight(..., trainable=True)` (the default). Gradients flow to them and the optimizer updates them each step. Listed in `layer.trainable_weights`. - **Non-trainable weights** — created with `add_weight(..., trainable=False)`. They are genuine variables with a dtype, a shape and a saved value, but no optimizer ever writes to them. Listed in `layer.non_trainable_weights`. `layer.weights` is the concatenation of the two, and `layer.count_params()` counts both. ## Why non-trainable variables exist Plenty of layers need memory that is *learned from data by a rule other than gradient descent*, or that is simply bookkeeping: - Running statistics — `BatchNormalization` accumulates `moving_mean` and `moving_variance` by exponential moving average during training and uses them at inference. - Counters — steps seen, tokens processed, calibration samples collected. - Frozen lookup content — an embedding table you deliberately do not train (though for that, `trainable=False` on the layer is the simpler route). - Quantization or normalization constants estimated during a calibration pass. ## Creating and updating them ``` def build(self, input_shape): self.total = self.add_weight( shape=(), initializer="zeros", trainable=False, name="total" ) ``` A `keras.Variable` exposes `assign(value)`, `assign_add(delta)` and `assign_sub(delta)`. Inside `call` you write `self.total.assign_add(keras.ops.sum(inputs))`. Reading it is just `self.total` used in ops, or `float(self.total)` / `keras.ops.convert_to_numpy(...)` outside the graph. Guard updates by mode where it matters: a running statistic should usually only move when `training` is true, which is why such layers declare `def call(self, inputs, training=None)`. ## Why a Python attribute is not a substitute Writing `self.total = 0.0` in `build` and `self.total += ...` in `call` looks equivalent and is not: 1. **Not tracked.** It never appears in `layer.weights`, so it is not part of the layer's state and is not written out when the model is saved or read back when it is loaded. Reload the model and your counter is whatever the constructor set. 2. **Not compatible with compiled execution.** Keras 3 runs on TensorFlow, JAX or PyTorch, and `call` is frequently traced into a compiled function rather than executed once per batch. Python-level side effects happen during tracing, not on every batch — so the number you see afterwards is meaningless. Real variables are the mechanism the backends use to express mutable state, and Keras threads their updates through for you (JAX's stateless requirement included). 3. **Not visible to tooling.** Summaries, weight counts and checkpoint diffing all work from the variable lists. ## The accounting subtlety of layer.trainable `trainable` on the *variable* and `trainable` on the *layer* interact. Setting `layer.trainable = False` does not convert your variables; it changes how the layer reports them — the layer's trainable weights are reported as empty and everything shows up under non-trainable weights, so the optimizer skips them. Flip it back to `True` and the original split returns. The variables themselves never changed; only the accounting did. This is why a frozen backbone shows a huge "non-trainable params" count in a model summary. ## Contrast with a buffer-free design If the value is a genuine constant — a fixed positional-encoding table, a scaling factor — you have a choice: store it as a non-trainable weight, or compute it in `call` with `keras.ops`. A weight costs checkpoint space but survives loading unchanged and is inspectable; a computed constant costs a little compute per call and cannot drift. For anything mutable, the variable is the only correct option. ## What an interviewer is checking That you know a layer's state is not just "the parameters", that you reach for `add_weight(trainable=False)` rather than a Python field, and that you can point at `BatchNormalization` as the built-in doing exactly this.
- Which built-in Keras layer is the canonical user of non-trainable weights?`BatchNormalization`. Its `moving_mean` and `moving_variance` are non-trainable variables updated by an exponential moving average during training and read at inference, while `gamma` and `beta` are the trainable pair. It is the clearest example of state learned from data by a rule that is not gradient descent.
- What does setting layer.trainable = False do to the layer's weight lists?It changes the reporting, not the variables: the layer's trainable weights come back empty and all of its weights are reported as non-trainable, so the optimizer skips them. Flipping it back restores the original split. The values are untouched throughout — which is why a frozen backbone shows a large non-trainable parameter count.
- Are non-trainable weights included when the model is saved?Yes. Saving persists the layer's full weight list, trainable and non-trainable alike, which is precisely why a reloaded model reproduces its inference-time normalization exactly. That is also the practical argument against keeping such state in a plain Python attribute: it would silently reset on load.
saying these in an interview costs you the question
- Keeping mutable layer state in a plain Python attribute
- Assuming non-trainable weights are excluded from saved models
- Thinking layer.trainable=False permanently rewrites each variable's flag
- Using += on a keras.Variable instead of assign_add
- Believing every entry in layer.weights receives gradients