Batch Normalization Cheat Sheet
Explains why batch normalization stabilizes training, its learnable scale and shift parameters, and how train versus eval mode statistics differ in PyTorch.
Core Concepts
What batch norm computes and why.
- Normalization step- For each mini-batch, subtract the batch mean and divide by the batch standard deviation, per feature/channel
- Learnable scale/shift (gamma, beta)- After normalizing, the layer applies y = gamma*x_hat + beta so it can undo normalization if that's optimal
- Internal covariate shift- The original motivation: reduce how much the distribution of layer inputs shifts as earlier layers update
- Running statistics- During training, exponential moving averages of batch mean/variance are tracked for use at inference
- Train vs. eval mode- Training uses current batch statistics; eval (model.eval()) uses the stored running mean/variance instead
Batch Normalization in PyTorch
Typical placement in a CNN block.
import torch.nn as nnblock = nn.Sequential( nn.Conv2d(64, 128, kernel_size=3, padding=1), nn.BatchNorm2d(128), # one gamma/beta per channel nn.ReLU(inplace=True),)# Critical: switch modes correctlymodel.train() # uses batch statistics, updates running statsmodel.eval() # uses stored running_mean / running_var, no updates
The Batch Norm Formula
What happens under the hood for a single feature.
# x: activations for one feature across the mini-batchmu = x.mean()var = x.var(unbiased=False)x_hat = (x - mu) / (var + eps).sqrt() # eps avoids divide-by-zero, e.g. 1e-5y = gamma * x_hat + beta # gamma, beta are learned per-feature
Practical Notes
Common pitfalls.
- Small batch sizes- Batch norm statistics get noisy with very small batches (e.g., < 8); consider GroupNorm or LayerNorm instead
- Forgetting model.eval()- The single most common inference bug -- leaves the model using batch statistics from the (often batch-size-1) inference input
- Placement relative to activation- Commonly Conv/Linear -> BatchNorm -> Activation, though BN-after-activation is also used in some architectures
- Bias term redundancy- The preceding Conv/Linear layer's bias is usually disabled (bias=False) since BatchNorm's beta already shifts the output
BatchNorm vs. Other Normalization Layers
How BatchNorm compares to alternatives that normalize over different axes.
- BatchNorm- Normalizes per-channel across the batch + spatial dims; statistics depend on other samples in the batch, so it struggles at batch size 1
- LayerNorm- Normalizes across the feature dimension for each sample independently; batch-size invariant, the default in transformers
- GroupNorm- Splits channels into groups and normalizes within each group per sample; a good BatchNorm replacement for small batches (e.g. detection/segmentation)
- InstanceNorm- Normalizes each channel of each sample independently (one group per channel); common in style transfer where batch statistics would leak content
- RMSNorm- Like LayerNorm but skips mean-centering, only rescaling by root-mean-square; cheaper and used in many modern LLMs
- Weight Standardization- Normalizes the convolution weights themselves rather than activations, often paired with GroupNorm
Fusing BatchNorm into a Preceding Conv for Inference
Folding BN's affine transform into the conv weights eliminates a memory-bound op and speeds up deployment.
import torch@torch.no_grad()def fuse_conv_bn(conv, bn): fused = torch.nn.Conv2d( conv.in_channels, conv.out_channels, conv.kernel_size, conv.stride, conv.padding, bias=True ) # scale = gamma / sqrt(running_var + eps) std = (bn.running_var + bn.eps).sqrt() scale = bn.weight / std fused.weight.copy_(conv.weight * scale.reshape(-1, 1, 1, 1)) bias = bn.bias - bn.running_mean * scale if conv.bias is not None: bias += conv.bias * scale fused.bias.copy_(bias) return fused# torch.ao.quantization.fuse_modules / torch.jit can automate this for whole models
SyncBatchNorm for Multi-GPU Training
Per-GPU mini-batches shrink effective batch statistics; SyncBatchNorm aggregates stats across all devices.
import torch.nn as nnfrom torch.nn.parallel import DistributedDataParallel as DDPmodel = build_model().cuda()# Convert every BatchNorm layer to SyncBatchNorm before wrapping in DDPmodel = nn.SyncBatchNorm.convert_sync_batchnorm(model)model = DDP(model, device_ids=[local_rank])# Without this, each GPU normalizes against only its own local shard,# which behaves like training with a much smaller batch size
Backpropagating Through BatchNorm by Hand
The gradient must flow through both the mean and variance, since every output depends on every input in the batch.
def batchnorm_backward(dout, x, x_hat, mu, var, gamma, eps): N = x.shape[0] std_inv = 1.0 / (var + eps).sqrt() dgamma = (dout * x_hat).sum(0) dbeta = dout.sum(0) dx_hat = dout * gamma dvar = (dx_hat * (x - mu) * -0.5 * std_inv**3).sum(0) dmu = (dx_hat * -std_inv).sum(0) + dvar * (-2.0 * (x - mu)).mean(0) dx = dx_hat * std_inv + dvar * 2 * (x - mu) / N + dmu / N return dx, dgamma, dbeta
Tuning momentum and eps
Two constructor arguments that are easy to leave at defaults but affect training stability.
- momentum default (0.1)- PyTorch's running stats update as running = (1-momentum)*running + momentum*batch_stat; lower momentum (e.g. 0.01) smooths noisy stats from small batches
- eps default (1e-5)- Raise eps (e.g. 1e-3) if you see NaNs from near-zero variance channels, common with mixed-precision training
- track_running_stats=False- Forces BatchNorm to always use current-batch statistics, even in eval() mode -- useful for meta-learning / few-shot setups
- Ghost Batch Norm- Splits a large batch into smaller virtual sub-batches for BN statistics, decoupling optimizer batch size from normalization batch size
- Cumulative moving average- Setting momentum=None makes PyTorch use a true cumulative average of all batches seen instead of an exponential moving average
Batch norm's running statistics are only updated during model.train() forward passes -- if you forget to call model.train() before a training loop, you'll silently train against stale statistics.