In Keras 3, what happens to a BatchNormalization layer when you set trainable=False?
answer
- one layer breaks the usual rule
- two of its four variables are not trained
- freezing changes behaviour, not just updates
- moving statistics stop moving
- recompile before the change bites
basics
~20 sBatchNormalization is special-cased: trainable=False also puts it in inference mode, so it normalizes with the stored moving_mean and moving_variance and stops updating them. For every other layer, trainable=False only stops the optimizer from updating weights.
solid answer
~50 s`BatchNormalization` holds four variables: trainable `gamma` and `beta`, and non-trainable `moving_mean` and `moving_variance` that are updated by a running average during training. For an ordinary layer, `layer.trainable = False` means only that the optimizer skips its weights. For `BatchNormalization` in Keras 3 it means more: the layer runs in inference mode, normalizing with the moving statistics and leaving them untouched. That special case exists because of transfer learning — if you freeze a pretrained backbone but its normalization layers keep re-estimating statistics on your small new dataset, the frozen convolution weights below are suddenly fed a different distribution and the pretrained features degrade. Two operational notes: a change to `trainable` takes effect on the training function only after you call `model.compile()` again, and when you later unfreeze the backbone for fine-tuning, people commonly still invoke it with `training=False` so the statistics stay frozen while the weights adapt.
code
python · 14 linesimport keras
base = keras.applications.ResNet50(include_top=False, weights=None)
base.trainable = False # also pins every BatchNormalization to inference mode
inputs = keras.Input(shape=(224, 224, 3))
x = base(inputs, training=False)
x = keras.layers.GlobalAveragePooling2D()(x)
outputs = keras.layers.Dense(1, activation="sigmoid")(x)
model = keras.Model(inputs, outputs)
model.compile(optimizer="adam", loss="binary_crossentropy")
bn = [l for l in base.layers if isinstance(l, keras.layers.BatchNormalization)][0]
print([w.name for w in bn.weights])go deeper
Know that a BatchNormalization layer stores running statistics used at prediction time, and that freezing a pretrained model is done with trainable = False.
Name the four variables and say which two the optimizer touches, and explain that this layer is special-cased so freezing also pins it to inference behaviour.
Walk a two-phase fine-tune end to end — freeze, compile, train the head, unfreeze, recompile at a low learning rate — and explain the train/evaluate gap that appears when statistics drift on a frozen backbone.
Weigh normalization choice as an architectural decision: batch-dependent statistics couple training batch size, fine-tuning procedure and serving behaviour, and sample-wise normalization removes that coupling at some cost in throughput and accuracy.
## What the layer actually holds `keras.layers.BatchNormalization()` owns four variables per feature channel: - `gamma` (scale) and `beta` (offset) — trainable, updated by the optimizer. - `moving_mean` and `moving_variance` — non-trainable, updated by an exponential moving average controlled by `momentum` (default 0.99) during training only. In training mode it normalizes the activations with the *current batch's* mean and variance, then updates the moving statistics toward those batch values. In inference mode it normalizes with the stored moving statistics and updates nothing. `epsilon` (default 1e-3) guards the division. ## The general meaning of trainable For any layer, `layer.trainable = False` tells Keras that the optimizer must not update that layer's weights: its variables are reported under `non_trainable_weights` and no gradient update is applied. Nothing about the layer's *computation* changes. Freeze a `Dense`, and it still multiplies by the same kernel it always did. ## The BatchNormalization special case Keras deliberately breaks that rule for `BatchNormalization`. When the layer is not trainable, it also runs in inference mode: it uses `moving_mean`/`moving_variance` to normalize and does not update them, regardless of whether the enclosing call is in training mode. This is a pragmatic decision, not an accident of implementation. Consider the standard transfer-learning recipe: take a pretrained backbone, set `base.trainable = False`, bolt a fresh head on top, and train the head on a few thousand images. If the backbone's normalization layers kept updating their statistics on the new data, then batch by batch the normalization applied between the frozen convolutions would drift. The frozen weights were tuned for the original statistics; changing what flows into them quietly destroys the very features you froze them to preserve. Worse, with small batches the batch estimates are noisy, so training and inference disagree. Freezing the statistics together with the weights is what "frozen backbone" is supposed to mean. ## The recompile requirement Setting `trainable` on layers of a model that has already been compiled does not retroactively change the compiled training function. The rule is: change `trainable`, then call `model.compile(...)` again before `fit`. Skip it and you get the confusing case where `model.summary()` shows the parameters as non-trainable while training still updates them. ## The two-phase fine-tune A typical schedule: 1. **Phase 1 — head only.** `base.trainable = False`, compile, train the new head at a normal learning rate. The base is frozen and its normalization statistics are pinned. 2. **Phase 2 — unfreeze.** `base.trainable = True`, recompile with a much smaller learning rate, continue training. In phase 2 the normalization layers become trainable again and would resume updating their statistics. The usual guard is to invoke the base with `training=False` inside a functional or subclassed forward pass — `x = base(inputs, training=False)` — so the statistics stay frozen while `gamma`, `beta` and the convolution weights adapt. Alternatively you can set `layer.trainable = False` on just the normalization layers while leaving the rest of the backbone trainable. ## Diagnosing it in the wild Symptoms that point here: training accuracy that looks fine while `evaluate` on the *same* data is much worse; a fine-tune whose first epoch immediately degrades a strong pretrained model; or metrics that shift when you change `batch_size` alone. All three are consistent with normalization statistics that do not match between the two modes. ## Related knobs `keras.layers.LayerNormalization` has no moving statistics — it normalizes each sample over its feature axes, identically in both modes — which is one reason architectures that are frequently fine-tuned or run at batch size 1 prefer it. If your batches are tiny, that difference matters more than the special case above. ## The one-sentence version `trainable=False` freezes weights everywhere; on `BatchNormalization` it also freezes the layer's *behaviour*, and remembering that distinction is the difference between a fine-tune that works and one that quietly ruins a pretrained backbone.
- Why does a change to layer.trainable require calling compile() again?`compile()` builds the training function, including which variables the optimizer will update. Flipping `trainable` afterwards changes the model's bookkeeping but not the already-built step, so the old behaviour persists. Recompiling rebuilds it. The tell-tale symptom is a summary that reports parameters as non-trainable while `fit` keeps changing them.
- During phase-two fine-tuning with the backbone unfrozen, how do people keep normalization statistics frozen?Invoke the backbone with the mode flag pinned — `x = base(inputs, training=False)` in the forward pass — so its `BatchNormalization` layers keep using the stored moving statistics even while `gamma`, `beta` and the convolution weights are updated. The alternative is to leave `trainable = False` on the normalization layers specifically.
- What does momentum control on a BatchNormalization layer?It is the exponential-moving-average coefficient for the moving mean and variance, defaulting to 0.99: each training step nudges the stored statistics a little toward the current batch's. A high value gives smooth, slow-adapting statistics; too high combined with a short training run leaves the moving statistics far from the true data distribution, so inference disagrees with training.
- Does LayerNormalization have the same freezing subtlety?No. `LayerNormalization` computes its statistics from each individual sample at every call, so it has no moving statistics and behaves identically in training and inference. Freezing it only stops its scale and offset from being updated — the ordinary meaning of `trainable=False`.
saying these in an interview costs you the question
- Assuming trainable=False only stops gradient updates, for every layer alike
- Thinking moving_mean and moving_variance are learned by the optimizer
- Freezing a backbone without recompiling and trusting the summary
- Believing BatchNormalization behaves identically in both modes
- Fine-tuning a frozen backbone while its normalization statistics keep drifting