What does the training argument in a Keras layer's call() do, and who sets it?
answer
- a per-call mode, not a global switch
- only some layers care about it
- who calls with True, who with False
- declared in the call signature
- forward it to nested layers
basics
~20 sIt selects train-time versus inference behaviour per call. Dropout drops units and BatchNormalization uses batch statistics only when it is true. fit() passes True, evaluate() and predict() pass False, and a bare model(x) defaults to inference.
solid answer
~40 sSome layers are mode-dependent: `Dropout` zeroes activations only during training, `BatchNormalization` normalizes with batch statistics and updates its moving averages only during training, and the random preprocessing layers augment only during training. Keras carries that mode as an explicit `training` argument threaded through `__call__`, not as a global mode switch on the model. `fit()` calls the model with `training=True`; `evaluate()` and `predict()` call it with `training=False`; calling `model(x)` yourself with nothing specified gives inference behaviour. In a custom layer you opt in by declaring it: `def call(self, inputs, training=None)`. Keras inspects the signature and only passes the flag to layers that accept it. Forward it explicitly to mode-dependent sub-layers — `self.dropout(x, training=training)` — and if you implement randomness by hand, branch on `training` yourself.
code
python · 19 linesimport keras
class DenseBlock(keras.layers.Layer):
def __init__(self, units, rate=0.5, **kwargs):
super().__init__(**kwargs)
self.dense = keras.layers.Dense(units, activation="relu")
self.norm = keras.layers.BatchNormalization()
self.dropout = keras.layers.Dropout(rate)
def call(self, inputs, training=None):
x = self.dense(inputs)
x = self.norm(x, training=training)
return self.dropout(x, training=training)
block = DenseBlock(16)
x = keras.ops.ones((4, 8))
print(keras.ops.sum(block(x, training=False) == block(x, training=False)))go deeper
Recall that Dropout and BatchNormalization behave differently while training, and that fit turns training mode on while predict turns it off. You should not have to set it by hand for standard workflows.
Explain that Keras only passes training to layers whose call declares it, and show the one-line habit of forwarding it into every mode-dependent sub-layer you compose.
Diagnose from symptoms: non-deterministic predictions, or a train/serve metric gap, traced back to a mode flag lost in a composite layer or a hand-written training loop.
Own the convention across a shared layer library — signatures that declare the flag, review rules for forwarding it, and determinism guarantees at serving so an inference path can never accidentally run in training mode.
## Why a flag exists at all A handful of layers are not pure functions of their input — they behave one way while learning and another way while predicting: - `keras.layers.Dropout(rate)` zeroes a random fraction of activations (and rescales the rest) during training, and is an identity function at inference. - `keras.layers.BatchNormalization()` normalizes using the current batch's mean/variance during training while updating `moving_mean` and `moving_variance`; at inference it normalizes with those stored moving statistics. - The random preprocessing/augmentation layers (for example `keras.layers.RandomFlip`, `keras.layers.RandomRotation`) transform during training and pass data through at inference. Everything else — `Dense`, `Conv2D`, `Embedding`, `LayerNormalization` — behaves identically either way. ## How Keras represents the mode Keras does not have a global "switch the model to eval" call. The mode is an argument on each invocation: `outputs = layer(inputs, training=True)`. That makes it explicit and local, and it means the same layer object can be used in both modes inside one step without mutating anything. Who supplies it: - `model.fit(...)` runs the forward pass with `training=True`. - `model.evaluate(...)` and `model.predict(...)` run it with `training=False`. - `model(x)` written by hand with no argument resolves to inference behaviour — `Dropout.call` itself defaults `training=False`. - You can always override: `model(x, training=True)` is how people sample dropout at prediction time for MC-dropout-style uncertainty. ## Declaring it in a custom layer Keras inspects your `call` signature. If it declares a `training` parameter, Keras passes the current value; if it does not, Keras passes nothing. So: ``` def call(self, inputs, training=None): ... ``` is how you opt in. A layer that has no mode-dependent behaviour should simply not declare the parameter. ## Forwarding to sub-layers The classic bug in a composite layer is: ``` def call(self, inputs, training=None): x = self.dense(inputs) return self.dropout(x) # training never forwarded ``` Write `self.dropout(x, training=training)` instead. Keras 3 does maintain a call context that lets a nested layer inherit the enclosing training value, so this often still behaves correctly — but the explicit pass-through is the documented pattern, it survives someone calling your layer standalone, and it makes the intent readable in review. Anything you implement with raw ops has no such fallback: if you write your own noise or your own normalization, you must branch on `training` yourself. For randomness, create a `keras.random.SeedGenerator` in `__init__` and pass it as the `seed` argument to `keras.random.*` calls so the layer stays reproducible and backend-agnostic. ## The symptoms when it goes wrong - **Dropout stuck on at inference**: predictions for the same input differ run to run, and evaluation metrics are worse than validation metrics computed during training. - **Dropout never on**: training loss falls suspiciously fast and the model overfits, because the regularizer you configured never fired. - **BatchNormalization in the wrong mode**: a large train/inference metric gap, especially with small batches, because the layer normalized with batch statistics it will not have at serving time. A quick check: call `model(x, training=False)` twice on identical input; if the outputs differ, something stochastic is still running. ## Custom training loops If you write your own step (with a gradient tape or the backend's autograd), you own the flag — the forward pass in the training step must pass `training=True`, and the evaluation pass must pass `training=False`. Nothing infers it for you outside `fit`/`evaluate`/`predict`. ## Not the same thing as trainable `training` is a per-call mode. `layer.trainable` is a persistent property controlling whether the optimizer updates that layer's weights. They are independent knobs — with the notable exception of `BatchNormalization`, which is deliberately special-cased so that freezing it also pins it to inference behaviour.
- If a custom layer's call() omits the training parameter entirely, what happens?Keras inspects the signature and simply never passes the flag to that layer, so the layer itself cannot branch on mode. Nested mode-dependent sub-layers can still pick up the enclosing value through Keras 3's call context, but any behaviour you implemented with raw ops inside that `call` will run identically in training and inference — usually silently wrong.
- How would you deliberately keep dropout active at prediction time?Call the model directly with the flag: `model(x, training=True)`, repeatedly, and aggregate the outputs. That is the Monte-Carlo dropout trick for a rough uncertainty estimate. `model.predict(x)` cannot do it — predict always runs the forward pass in inference mode.
- In a custom training loop, what sets the training flag?You do. Outside `fit`/`evaluate`/`predict` nothing infers it, so your training step must call `model(x, training=True)` and your validation step `model(x, training=False)`. Getting this backwards is a common cause of a custom loop whose numbers do not match `fit`'s.
saying these in an interview costs you the question
- Assuming Keras has a global eval mode like a module switch
- Forgetting to forward training into nested Dropout or BatchNormalization
- Believing Dense and Conv2D behave differently in training mode
- Thinking model.predict() can run dropout stochastically
- Confusing the training call flag with the persistent trainable property