Course Content
Deep Learning with TensorFlow and PyTorch
4 sections · 15 lessons
Dropout, Batch Normalization, and Callbacks
A model reaches 94% validation accuracy in training. You save it, load it in a small service, send it a single image, and get a confident, wrong answer. Send the same image again inside a batch of 32 other images and it comes back correct.
The same input. Two different predictions. Depending on what its neighbours were.
That is not a bug in your serving code. It is batch normalisation left in training mode. In that mode, a batch-norm layer normalises each feature using the mean and variance of the current batch, so every prediction depends on the other examples that happened to travel with it. With a batch of one, the statistics come from that single example alone: in a dense layer each value is its own mean, so every feature normalises to zero whatever the input (PyTorch refuses outright with "Expected more than 1 value per channel when training"), and in a convolutional layer the image is normalised against itself, which matches nothing the model saw in training.
Batch normalisation and dropout are the two most-used layers in deep learning that behave differently at training time and inference time. Both are worth understanding at the level of the arithmetic, because almost every problem they cause comes from a mismatch between those two modes.
What batch normalisation actually computes
For each feature (or channel) independently, over the examples in the current batch:
Take one feature with four values in the batch: [2,4,6,8].
| Step | Computation | Result |
|---|---|---|
| Batch mean | (2+4+6+8)/4 | 5.0 |
| Batch variance | (9+1+1+9)/4 | 5.0 |
| Normalise | (x−5)/5.0+10−5 | [−1.342, −0.447, 0.447, 1.342] |
| Scale and shift (γ=2, β=1) | 2x^+1 | [−1.683, 0.106, 1.894, 3.683] |
The γ and β are learned parameters, two per feature, and they are what stops this from being a straitjacket. Forcing every layer's output to have mean 0 and variance 1 would destroy information the network might need; γ and β let the network undo the normalisation if that turns out to be the right thing to do. In the extreme, setting γ=σB and β=μB reproduces the original values exactly. The layer can be the identity; it just does not start there.
The ϵ (typically 10−5) is not cosmetic. Without it, a feature that happens to be constant across a batch has zero variance, and the division becomes zero divided by zero — nan.
The two modes, and the statistics that bridge them
At inference you cannot use batch statistics — you may have a batch of one, and in any case predictions must not depend on unrelated inputs. So during training the layer also maintains an exponential moving average of the statistics it has seen:
model.train() # BN uses THIS BATCH's mean/var, and updates the running averagesmodel.eval() # BN uses the stored running averages, updates nothingOne cross-framework trap: the two libraries define the momentum parameter in opposite directions.
PyTorch momentum | Keras momentum | |
|---|---|---|
| Default | 0.1 | 0.99 |
| Meaning | Weight given to the new batch | Weight given to the old running value |
| Equivalent setting | 0.1 | 0.9 |
Copy a momentum of 0.99 from a Keras model into PyTorch and your running statistics will be dominated by the most recent batch, making inference behaviour erratic. This is a genuine source of "the port doesn't match the original" bugs.
A related failure: if you train with very few steps, or freeze a pretrained backbone in train() mode on a tiny fine-tuning set, the running statistics can be badly estimated or drift to values that suit your small dataset rather than the original one. When fine-tuning a pretrained network on a few hundred images, it is common practice to keep the batch-norm layers in eval() mode throughout, so they use the well-estimated statistics from the original large-scale training.
1model.train()2for m in model.modules():3 if isinstance(m, nn.BatchNorm2d):4 m.eval() # freeze BN statistics while fine-tuningWhy it helps — the honest version
The original paper attributed the gain to reducing "internal covariate shift": the idea that as earlier layers update, the distribution of inputs to later layers keeps moving, so later layers spend their capacity chasing a moving target. It is an appealing story and it is at best incomplete — later work injected deliberate distribution shift after batch norm layers and found training still improved.
The better-supported explanation is that batch norm smooths the loss landscape. By constraining the scale of each layer's outputs, it bounds how much the loss and its gradients can change for a given step, which means larger learning rates remain stable. That, empirically, is where most of the speed-up comes from.
What it reliably delivers:
| Benefit | Practical effect |
|---|---|
| Tolerates much larger learning rates | Often 5–10× larger, so far fewer epochs to converge |
| Reduces sensitivity to initialisation | A merely reasonable init works where it previously would not |
| Mild regularisation | Batch statistics are noisy, which acts like a small amount of noise injection |
| Keeps activations in a useful range | Fewer saturated units, fewer dead ReLUs |
Placement, and the bias that becomes useless
The conventional order is Linear → BatchNorm → Activation. Normalising the pre-activation keeps values in the region where the activation function has a healthy gradient, which is the point.
1nn.Sequential(2 nn.Linear(256, 128, bias=False), # bias is redundant -- see below3 nn.BatchNorm1d(128),4 nn.ReLU(),5)bias=False is correct and worth understanding. The linear layer's bias adds a constant to every value in a feature; batch norm then subtracts the batch mean, which removes that constant entirely. The bias has no effect on the output and merely wastes parameters and a small amount of compute. Batch norm's own β plays the role of the bias.
Some architectures put normalisation before the linear layer instead (pre-activation residual blocks, and the pre-norm arrangement standard in transformers), which improves gradient flow in very deep stacks. Both orders work; follow the convention of the architecture family you are implementing.
The normalisation family
Batch norm's dependence on batch statistics is its weakness. If your batch is 4 images because each one is 512×512, the mean and variance estimates are noisy, and performance degrades. The alternatives normalise over different axes and are therefore batch-size independent.
| Layer | Normalises over | Depends on batch size? | Standard use |
|---|---|---|---|
| BatchNorm | The batch, per channel | Yes — needs 16+ | CNNs with reasonable batch sizes |
| LayerNorm | All features, per example | No | Transformers, RNNs — the default in NLP |
| InstanceNorm | Spatial dims, per example per channel | No | Style transfer, image generation |
| GroupNorm | Groups of channels, per example | No | Detection and segmentation, where batches are small |
1nn.BatchNorm2d(64) # per-channel, across the batch2nn.LayerNorm(512) # across the feature dimension of each example3nn.GroupNorm(num_groups=8, num_channels=64)4nn.InstanceNorm2d(64)A useful decision rule: if your batch size per device is below about 16, do not use batch norm. Use GroupNorm for vision, LayerNorm for sequences. The failure mode of small-batch batch norm is subtle — training looks fine, and evaluation is unaccountably worse — so it costs a lot of debugging time to discover the hard way.
Dropout, at the level of the mask
During training, dropout samples a Bernoulli mask and applies it element-wise, then rescales:
The division by (1−p) is inverted dropout, and it exists so that inference needs no adjustment at all. Consider E[yi]: with probability 1−p the value survives and is scaled by 1/(1−p); with probability p it is zero. The expectation is (1−p)⋅xi/(1−p)=xi. Scale preserved.
1import torch23x = torch.tensor([[2., 4., 6., 8.]])4drop = torch.nn.Dropout(p=0.5)56drop.train()7print(drop(x)) # e.g. tensor([[ 0., 8., 0., 16.]]) -- survivors doubled8drop.eval()9print(drop(x)) # tensor([[2., 4., 6., 8.]]) -- identityTwo consequences follow directly. Training loss is measured with units randomly missing, so it is higher than the model's true training loss — which is why validation loss sometimes appears better than training loss early on, and why that is not a bug. And in eval() mode dropout is exactly the identity function, doing nothing but costing a function call.
Dropout that respects structure
Standard dropout zeroes individual elements independently, which is the right thing for a dense layer and the wrong thing in two common cases.
Convolutional feature maps. Neighbouring pixels in a feature map are highly correlated, so zeroing individual pixels removes almost no information — the surrounding pixels carry it. nn.Dropout2d drops entire channels instead, which actually removes a feature.
Recurrent networks. Applying a fresh random mask at every timestep injects noise that compounds across a long sequence and destroys the recurrent state. Variational dropout samples one mask and reuses it at every timestep for a given sequence, which regularises without destroying memory. This is what PyTorch's dropout= argument to nn.LSTM approximates, and it applies only between layers, not within the recurrence.
nn.Dropout2d(0.1) # drops whole channelsnn.LSTM(input_size=128, hidden_size=256, num_layers=2, dropout=0.3)Monte Carlo dropout: uncertainty for free
Leaving dropout on at inference and running the same input many times gives a distribution of predictions rather than a point estimate. The spread is a usable proxy for model uncertainty.
1def mc_predict(model, x, n=50):2 model.eval()3 for m in model.modules():4 if isinstance(m, torch.nn.Dropout):5 m.train() # dropout ON, batchnorm still OFF6 with torch.no_grad():7 preds = torch.stack([model(x) for _ in range(n)])8 return preds.mean(0), preds.std(0) # prediction and its uncertaintyNote carefully that only the dropout modules are switched back on. Turning the whole model to train() would also re-enable batch-norm batch statistics, reintroducing exactly the bug this lesson opened with.
Why combining them can backfire
Dropout and batch norm together are known to underperform, and the mechanism is a variance mismatch. During training, dropout's random masking increases the variance of the activations reaching the next batch-norm layer, and the running statistics are estimated from that inflated variance. At inference, dropout is off, so the variance drops — but batch norm is still normalising with statistics calibrated to the noisier training-time distribution. The layer's outputs are systematically mis-scaled.
| Situation | Recommendation |
|---|---|
| Convolutional network with batch norm | Skip dropout entirely; BN plus augmentation plus weight decay is enough |
| Dense network with batch norm | Dropout after the activation, and at a lower rate (0.1–0.2) |
| Transformer | LayerNorm plus dropout — this combination is fine, because LayerNorm does not use batch statistics |
| No normalisation layers | Dropout at the usual 0.3–0.5 |
Callbacks: automating the decisions you would otherwise make by hand
A callback is a piece of code registered to run at a defined point in training — the end of a batch, the end of an epoch, the start of training. It exists so that logic which is not part of the model can still participate in the loop: saving checkpoints, adjusting the learning rate, stopping early, logging.
1from tensorflow import keras23callbacks = [4 keras.callbacks.EarlyStopping(5 monitor="val_loss", patience=10, restore_best_weights=True),6 keras.callbacks.ModelCheckpoint(7 "best.keras", monitor="val_loss", save_best_only=True),8 keras.callbacks.ReduceLROnPlateau(9 monitor="val_loss", factor=0.5, patience=4, min_lr=1e-6),10 keras.callbacks.TensorBoard(log_dir="logs/run1", histogram_freq=1),11 keras.callbacks.CSVLogger("history.csv"),12]1314model.fit(train_ds, validation_data=val_ds, epochs=200, callbacks=callbacks)Note the interaction between the first and third: ReduceLROnPlateau needs a shorter patience than EarlyStopping, or training will stop before the learning rate ever gets cut. A common pairing is patience 4 for the LR reduction and 10 for the stop, giving the reduced rate two chances to produce an improvement.
Writing your own is a matter of overriding the hook you care about:
1import numpy as np23class DivergenceGuard(keras.callbacks.Callback):4 """Warn when the loss stops being finite -- catches divergence immediately."""5 def on_batch_end(self, batch, logs=None):6 loss = (logs or {}).get("loss")7 if loss is not None and not np.isfinite(loss):8 print(f"\nNon-finite loss at batch {batch}; stopping.")9 self.model.stop_training = True1011 def on_epoch_end(self, epoch, logs=None):12 gap = logs["val_loss"] - logs["loss"]13 print(f" overfit gap: {gap:+.4f}")PyTorch has no callback system, because you own the loop. The equivalent is simply code in the right place:
1best_val, wait, patience = float("inf"), 0, 102sched = torch.optim.lr_scheduler.ReduceLROnPlateau(opt, factor=0.5, patience=4)34for epoch in range(200):5 train_loss = train_one_epoch(model, opt, train_loader)6 val_loss = evaluate(model, val_loader)78 sched.step(val_loss) # LR reduction910 if val_loss < best_val - 1e-4: # checkpointing11 best_val, wait = val_loss, 012 torch.save({"epoch": epoch,13 "model": model.state_dict(),14 "opt": opt.state_dict(),15 "val_loss": val_loss}, "best.pt")16 else:17 wait += 118 if wait >= patience: # early stopping19 break2021 writer.add_scalar("loss/train", train_loss, epoch) # logging22 writer.add_scalar("loss/val", val_loss, epoch)Six lines replace four callbacks. The trade is explicitness against boilerplate, which is the same trade the two frameworks make everywhere.
The checks that prevent all of this
Almost every problem in this lesson is a train-versus-inference mismatch, and two habits catch nearly all of them.
Evaluate the way you will deploy. After training, load the saved weights into a fresh process, call model.eval(), and run a single example — batch size one, exactly as production will. If the number differs from your validation score by more than rounding, you have a mode bug. Finding it here costs ten minutes; finding it after release costs considerably more.
Assert the mode rather than trusting it. A one-line check at the top of your evaluation function — assert not model.training — has prevented more bad numbers than any amount of care. The failure it catches is silent: the model runs, produces plausible predictions, and is simply wrong by a few percent in a direction you cannot see.
Beyond that, keep the decision rules straight. Batch size under 16 means GroupNorm or LayerNorm, not BatchNorm. Batch norm present in a convolutional network means you probably do not need dropout at all. Fine-tuning a pretrained model on a small dataset means freezing the batch-norm statistics. And any layer whose behaviour differs between training and inference deserves a moment's thought about which mode it is in every time you write an inference path.