Deep Learning with TensorFlow and PyTorch

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.

Why a batch of one gives a different answertrain() mode• BatchNorm uses thisbatch's mean and variance• Those statistics feed a running average• Dropout zeroes a fresh random mask• Correct for learning, wrong for servingeval() mode• BatchNorm uses thestored running statistics• Result no longer depends on batch mates• Dropout is off and nothing is scaled• One image gives the same answer as 32
With a batch of one, BatchNorm normalises the example against itself, so a dense layer's features all come out as zero whatever the input.

What batch normalisation actually computes

For each feature (or channel) independently, over the examples in the current batch:

μB=1m∑i=1mxi,σB2=1m∑i=1m(xi−μB)2\mu_B = \frac{1}{m}\sum_{i=1}^{m} x_i, \qquad \sigma_B^2 = \frac{1}{m}\sum_{i=1}^{m}(x_i - \mu_B)^2

x^i=xi−μBσB2+ϵ,yi=γx^i+β\hat{x}_i = \frac{x_i - \mu_B}{\sqrt{\sigma_B^2 + \epsilon}}, \qquad y_i = \gamma \hat{x}_i + \beta

Take one feature with four values in the batch: [2,4,6,8][2, 4, 6, 8].

StepComputationResult
Batch mean(2+4+6+8)/4(2+4+6+8)/45.0
Batch variance(9+1+1+9)/4(9+1+1+9)/45.0
Normalise(x−5)/5.0+10−5(x - 5)/\sqrt{5.0 + 10^{-5}}[−1.342, −0.447, 0.447, 1.342]
Scale and shift (γ=2\gamma=2, β=1\beta=1)2x^+12\hat{x} + 1[−1.683, 0.106, 1.894, 3.683]

The γ\gamma and β\beta 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; γ\gamma and β\beta let the network undo the normalisation if that turns out to be the right thing to do. In the extreme, setting γ=σB\gamma = \sigma_B and β=μB\beta = \mu_B reproduces the original values exactly. The layer can be the identity; it just does not start there.

The ϵ\epsilon (typically 10−510^{-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:

μrunning←(1−α) μrunning+α μB\mu_{\text{running}} \leftarrow (1-\alpha)\,\mu_{\text{running}} + \alpha\,\mu_B

Python
model.train()   # BN uses THIS BATCH's mean/var, and updates the running averagesmodel.eval()    # BN uses the stored running averages, updates nothing

One cross-framework trap: the two libraries define the momentum parameter in opposite directions.

PyTorch momentumKeras momentum
Default0.10.99
MeaningWeight given to the new batchWeight given to the old running value
Equivalent setting0.10.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.

Python
model.train()for m in model.modules():    if isinstance(m, nn.BatchNorm2d):        m.eval()          # freeze BN statistics while fine-tuning

Why 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:

BenefitPractical effect
Tolerates much larger learning ratesOften 5–10× larger, so far fewer epochs to converge
Reduces sensitivity to initialisationA merely reasonable init works where it previously would not
Mild regularisationBatch statistics are noisy, which acts like a small amount of noise injection
Keeps activations in a useful rangeFewer 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.

Python
nn.Sequential(    nn.Linear(256, 128, bias=False),   # bias is redundant -- see below    nn.BatchNorm1d(128),    nn.ReLU(),)

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 β\beta 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.

LayerNormalises overDepends on batch size?Standard use
BatchNormThe batch, per channelYes — needs 16+CNNs with reasonable batch sizes
LayerNormAll features, per exampleNoTransformers, RNNs — the default in NLP
InstanceNormSpatial dims, per example per channelNoStyle transfer, image generation
GroupNormGroups of channels, per exampleNoDetection and segmentation, where batches are small
Python
nn.BatchNorm2d(64)              # per-channel, across the batchnn.LayerNorm(512)               # across the feature dimension of each examplenn.GroupNorm(num_groups=8, num_channels=64)nn.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:

mi∼Bernoulli(1−p),yi=mi xi1−pm_i \sim \text{Bernoulli}(1-p), \qquad y_i = \frac{m_i \, x_i}{1-p}

The division by (1−p)(1-p) is inverted dropout, and it exists so that inference needs no adjustment at all. Consider E[yi]\mathbb{E}[y_i]: with probability 1−p1-p the value survives and is scaled by 1/(1−p)1/(1-p); with probability pp it is zero. The expectation is (1−p)⋅xi/(1−p)=xi(1-p) \cdot x_i/(1-p) = x_i. Scale preserved.

Python
import torchx = torch.tensor([[2., 4., 6., 8.]])drop = torch.nn.Dropout(p=0.5)drop.train()print(drop(x))     # e.g. tensor([[ 0.,  8.,  0., 16.]])  -- survivors doubleddrop.eval()print(drop(x))     # tensor([[2., 4., 6., 8.]])           -- identity

Two 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.

Python
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.

Python
def mc_predict(model, x, n=50):    model.eval()    for m in model.modules():        if isinstance(m, torch.nn.Dropout):            m.train()                      # dropout ON, batchnorm still OFF    with torch.no_grad():        preds = torch.stack([model(x) for _ in range(n)])    return preds.mean(0), preds.std(0)     # prediction and its uncertainty

Note 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.

SituationRecommendation
Convolutional network with batch normSkip dropout entirely; BN plus augmentation plus weight decay is enough
Dense network with batch normDropout after the activation, and at a lower rate (0.1–0.2)
TransformerLayerNorm plus dropout — this combination is fine, because LayerNorm does not use batch statistics
No normalisation layersDropout 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.

Python
from tensorflow import kerascallbacks = [    keras.callbacks.EarlyStopping(        monitor="val_loss", patience=10, restore_best_weights=True),    keras.callbacks.ModelCheckpoint(        "best.keras", monitor="val_loss", save_best_only=True),    keras.callbacks.ReduceLROnPlateau(        monitor="val_loss", factor=0.5, patience=4, min_lr=1e-6),    keras.callbacks.TensorBoard(log_dir="logs/run1", histogram_freq=1),    keras.callbacks.CSVLogger("history.csv"),]model.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:

Python
import numpy as npclass DivergenceGuard(keras.callbacks.Callback):    """Warn when the loss stops being finite -- catches divergence immediately."""    def on_batch_end(self, batch, logs=None):        loss = (logs or {}).get("loss")        if loss is not None and not np.isfinite(loss):            print(f"\nNon-finite loss at batch {batch}; stopping.")            self.model.stop_training = True    def on_epoch_end(self, epoch, logs=None):        gap = logs["val_loss"] - logs["loss"]        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:

Python
best_val, wait, patience = float("inf"), 0, 10sched = torch.optim.lr_scheduler.ReduceLROnPlateau(opt, factor=0.5, patience=4)for epoch in range(200):    train_loss = train_one_epoch(model, opt, train_loader)    val_loss = evaluate(model, val_loader)    sched.step(val_loss)                                  # LR reduction    if val_loss < best_val - 1e-4:                        # checkpointing        best_val, wait = val_loss, 0        torch.save({"epoch": epoch,                    "model": model.state_dict(),                    "opt": opt.state_dict(),                    "val_loss": val_loss}, "best.pt")    else:        wait += 1        if wait >= patience:                              # early stopping            break    writer.add_scalar("loss/train", train_loss, epoch)    # logging    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.