What Is Batch Normalization and How Does It Stabilize Neural Network Training?
Batch Normalization (BN) is a technique that normalizes layer activations to zero-mean and unit-variance for each mini-batch, which mitigates internal covariate shift and allows neural networks to train with higher learning rates and faster convergence.
Batch Normalization is a critical optimization technique for training deep neural networks, originally introduced to address the problem of changing input distributions between layers. In the karpathy/nn-zero-to-hero repository, Andrej Karpathy implements a pedagogical from-scratch version in the Makemore series to demonstrate exactly how BN stabilizes training. This article examines the mechanics of Batch Normalization using the actual source code from lectures/makemore/makemore_part3_bn.ipynb.
The Mathematical Foundation of Batch Normalization
Batch Normalization operates on activation tensors by enforcing a zero-mean, unit-variance distribution before the non-linearity is applied. For a given activation tensor h with shape [batch, dim], the algorithm performs four core steps.
Step 1: Compute Batch Statistics
First, calculate the mean and variance across the batch dimension.
μ = h.mean(0, keepdim=True) # batch mean (shape: [1, dim])
σ² = h.var(0, keepdim=True) # batch variance (shape: [1, dim])
Step 2: Normalize
Subtract the mean and divide by the standard deviation, adding a small epsilon for numerical stability.
ĥ = (h - μ) / torch.sqrt(σ² + eps)
Step 3: Scale and Shift
Apply learnable parameters gamma (γ) and beta (β) to allow the network to recover any required distribution.
out = γ * ĥ + β
Step 4: Maintain Running Statistics
During training, maintain a moving average of the batch statistics using a momentum term. These running statistics replace the batch statistics during inference when training=False.
From-Scratch Implementation in nn-zero-to-hero
The BatchNorm1d class in lectures/makemore/makemore_part3_bn.ipynb (lines 400-430) implements these steps manually to reveal the underlying mechanism. This custom implementation mirrors the behavior of torch.nn.BatchNorm1d used later in the notebook (lines 540-570).
import torch
class BatchNorm1d:
def __init__(self, dim, eps=1e-5, momentum=0.1):
self.eps = eps
self.momentum = momentum
self.training = True
# learnable parameters
self.gamma = torch.ones(dim)
self.beta = torch.zeros(dim)
# running statistics
self.running_mean = torch.zeros(dim)
self.running_var = torch.ones(dim)
def __call__(self, x):
if self.training:
mean = x.mean(0, keepdim=True)
var = x.var(0, keepdim=True)
else:
mean = self.running_mean
var = self.running_var
x_hat = (x - mean) / torch.sqrt(var + self.eps)
out = self.gamma * x_hat + self.beta
if self.training:
# update running stats with momentum
self.running_mean = (1 - self.momentum) * self.running_mean + self.momentum * mean
self.running_var = (1 - self.momentum) * self.running_var + self.momentum * var
return out
How Batch Normalization Stabilizes Training
According to the nn-zero-to-hero source code and empirical experiments, BN stabilizes training through four primary mechanisms.
Reduces internal covariate shift. By keeping the distribution of layer inputs stable, downstream layers see more predictable data, making gradient updates less noisy and preventing the distribution of inputs to deep layers from shifting drastically during training.
Allows higher learning rates. Normalized activations prevent exploding or vanishing gradients, enabling the optimizer to safely take larger steps and accelerating convergence without causing divergence.
Acts as a regularizer. The stochasticity introduced by computing statistics over mini-batches adds noise to the activations, functioning similarly to dropout and improving generalization.
Improves gradient flow. Normalization reduces the chance of saturated activations (e.g., in tanh or sigmoid layers), keeping gradients larger and more consistent across deep stacks of layers.
Empirically, the notebook demonstrates that adding BN dramatically lowers the final loss from approximately 2.12 to 2.07 and produces smoother training curves (lines 560-590), confirming its stabilizing effect.
Integrating BatchNorm into Training Pipelines
When using the custom BatchNorm1d class or PyTorch's built-in equivalent, you must distinguish between training and inference modes.
Building a Network with Custom BN
n_embd, n_hidden = 10, 100
block_size = 3
layers = [
torch.nn.Linear(n_embd * block_size, n_hidden, bias=False),
BatchNorm1d(n_hidden), # custom implementation
torch.nn.Tanh(),
# ... additional layers
]
# forward pass
h = emb.view(emb.shape[0], -1) # flatten
for layer in layers:
h = layer(h)
Using PyTorch's Built-in BatchNorm
For production models, use the optimized CUDA kernels in PyTorch.
import torch.nn as nn
layers = [
nn.Linear(n_embd * block_size, n_hidden, bias=False),
nn.BatchNorm1d(n_hidden), # built-in BN
nn.Tanh(),
]
model = nn.Sequential(*layers)
# Toggle modes
model.train() # uses batch statistics, updates running_mean/var
model.eval() # uses running_mean/var for inference
Monitoring Activation Distributions
As shown in the notebook, you can verify BN's effect by inspecting activations after the non-linearity.
# After forward pass through Tanh layers
for i, layer in enumerate(layers):
if isinstance(layer, nn.Tanh):
t = layer.out # stored during forward hook
print(f'layer {i}: mean {t.mean():+.2f}, std {t.std():.2f}')
Summary
- Batch Normalization normalizes activations to zero-mean and unit-variance using batch statistics during training.
- The technique is implemented from scratch in
karpathy/nn-zero-to-herovia theBatchNorm1dclass inlectures/makemore/makemore_part3_bn.ipynb. - Learnable parameters
gammaandbetaallow the network to denormalize activations, while running statistics cached via momentum enable stable inference. - BN stabilizes training by reducing internal covariate shift, permitting higher learning rates, acting as a regularizer, and improving gradient flow through deep networks.
- Empirical results show BN reduces loss significantly (from ~2.12 to ~2.07) and smooths training curves.
Frequently Asked Questions
What is internal covariate shift?
Internal covariate shift refers to the phenomenon where the distribution of inputs to a layer changes during training as the parameters of previous layers update. This forces each layer to continuously adapt to new input distributions, potentially slowing training and requiring lower learning rates. Batch Normalization mitigates this by constraining the mean and variance of layer inputs to remain stable.
Why are gamma and beta parameters necessary if we are normalizing?
The normalization step forces activations into a standard normal distribution (zero mean, unit variance), which may not be optimal for the subsequent layer or the non-linearity. The learnable gamma (scale) and beta (shift) parameters allow the network to undo the normalization if necessary, restoring representational power. Without these parameters, the network might lose the ability to model identity functions or preserve important feature scales.
How does Batch Normalization behave differently during training versus inference?
During training, BN computes statistics (mean and variance) from the current mini-batch and updates running averages using the momentum parameter. During inference (training=False), BN uses the cached running statistics rather than batch statistics to ensure deterministic outputs regardless of batch size. This requires explicitly setting the module to eval mode using model.eval() in PyTorch.
Can Batch Normalization be applied to convolutional layers?
Yes, for convolutional layers you use BatchNorm2d (or BatchNorm1d for sequences), which normalizes across the channel dimension and spatial locations. In nn-zero-to-hero, the makemore_part5_cnn1.ipynb notebook demonstrates using torch.nn.BatchNorm2d in convolutional networks, where it computes mean and variance across the batch and spatial dimensions (N, H, W) for each channel independently.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →