“What is the difference between batch norm and layer norm?” is one of the most common deep learning interview questions because a one-line answer is easy and a good answer is not. Interviewers want to hear which axis each layer averages over, then push on the consequences: why a model is accurate in training and poor in production, why fine-tuning with batch size 2 falls apart, and why transformers use LayerNorm or RMSNorm. If you can derive those consequences from the axis, you understand the layer.
Before you start
You should know what a mini-batch is, what an activation is, and how mean and variance are computed. The examples are pure Python 3.12 or newer with the standard library. The PyTorch snippets are illustrative (PyTorch 2.x) and were not executed for this article. Throughout, a batch is a matrix with one row per sample and one column per feature.
The short answer
Both layers subtract a mean, divide by a standard deviation, and then apply a learned per-feature scale (gamma) and shift (beta). Batch norm computes the mean and variance of each feature across the samples in the batch, so it needs reasonably large batches and keeps running averages for use at inference; it behaves differently in train() and eval() mode. Layer norm computes the mean and variance of each sample across its own features, so it is identical at training and inference time and does not care about batch size. Transformers use layer norm (or its cheaper cousin RMSNorm) because batches of variable-length sequences make batch statistics awkward and autoregressive inference often runs one sequence at a time.
How it works
Normalization keeps the inputs to each layer on a stable scale while the weights below it keep changing. The original batch norm paper credited reduced “internal covariate shift”; later work pointed to a smoother loss landscape that permits higher learning rates. The arithmetic is simple, and the only real difference between the two layers is the direction you average in:
import math
def mean(xs):
return sum(xs) / len(xs)
def var(xs):
m = mean(xs)
return sum((x - m) ** 2 for x in xs) / len(xs) # biased (divide by n), as the layers do
def batch_norm(X, eps=1e-5):
stats = [(mean(col), var(col)) for col in zip(*X)] # one (mean, var) per feature
return [[(x - m) / math.sqrt(v + eps) for x, (m, v) in zip(row, stats)] for row in X]
def layer_norm(X, eps=1e-5):
out = []
for row in X: # one (mean, var) per sample
m, v = mean(row), var(row)
out.append([(x - m) / math.sqrt(v + eps) for x in row])
return out
def show(M):
for row in M:
print([round(x, 3) for x in row])
X = [[1.0, 2.0, 3.0, 4.0], # sample 0
[2.0, 4.0, 6.0, 8.0], # sample 1
[3.0, 0.0, 9.0, 0.0]] # sample 2The small eps guards against a zero variance. In a CNN, BatchNorm2d keeps one mean and variance per channel and averages over the batch and both spatial dimensions. In a transformer, LayerNorm(d_model) keeps one mean and variance per token and averages over the hidden dimension.
Step-by-step walkthrough
Step 1: Batch norm, column by column
Take feature 1, the second column. Its values across the batch are 2, 4 and 0, with mean 2 and variance 8/3. Each value is replaced by its distance from that mean in standard deviations:
col = [row[1] for row in X]
print(col, mean(col), round(var(col), 4)) # [2.0, 4.0, 0.0] 2.0 2.6667
show(batch_norm(X))
# [-1.225, 0.0, -1.225, 0.0]
# [0.0, 1.225, 0.0, 1.225]
# [1.225, -1.225, 1.225, -1.225]
print([(round(mean(c), 6), round(var(c), 4)) for c in zip(*batch_norm(X))])
# [(0.0, 1.0), (0.0, 1.0), (0.0, 1.0), (0.0, 1.0)]Every column now has mean 0 and variance 1. The cost: sample 2’s output now depends on samples 0 and 1, because it describes where sample 2 sits relative to its batch-mates.
Step 2: Layer norm, row by row
show(layer_norm(X))
# [-1.342, -0.447, 0.447, 1.342]
# [-1.342, -0.447, 0.447, 1.342]
# [0.0, -0.816, 1.633, -0.816]
print(layer_norm([X[2]])[0] == layer_norm(X)[2]) # True
print(batch_norm([X[2], X[0]])[0] == batch_norm(X)[2]) # FalseEach row now has mean 0 and variance 1. Samples 0 and 1 produce the same output (to three decimals; only eps separates them) because sample 1 is sample 0 doubled, and layer norm removes per-sample scale. The two print lines show the property that matters most: a sample’s layer norm output is the same whether it is alone or in a batch, while its batch norm output changes when the batch changes.
Step 3: Train mode versus eval mode
At inference you may get one request at a time, so batch norm cannot rely on batch statistics. During training it keeps exponential moving averages, and in eval mode it uses those frozen values instead. This class follows PyTorch’s conventions: momentum=0.1 is the weight of the new batch, and the running variance uses the unbiased estimate (divide by n - 1) even though the batch is normalized with the biased one.
import random
class BatchNorm1d:
"""A single-feature batch norm that follows PyTorch's conventions."""
def __init__(self, momentum=0.1, eps=1e-5):
self.momentum, self.eps = momentum, eps
self.running_mean, self.running_var = 0.0, 1.0
self.gamma, self.beta = 1.0, 0.0
self.training = True
def __call__(self, xs):
if self.training:
m, v = mean(xs), var(xs) # normalize with this batch's statistics
n = len(xs)
self.running_mean = (1 - self.momentum) * self.running_mean + self.momentum * m
self.running_var = (1 - self.momentum) * self.running_var + self.momentum * v * n / (n - 1)
else:
m, v = self.running_mean, self.running_var # eval: frozen estimates
return [self.gamma * (x - m) / math.sqrt(v + self.eps) + self.beta for x in xs]
rng = random.Random(0)
bn = BatchNorm1d()
for step in range(200):
bn([rng.gauss(50.0, 10.0) for _ in range(32)]) # true mean 50, true variance 100
print(round(bn.running_mean, 2), round(bn.running_var, 2)) # 50.48 90.08
bn.training = False
print([round(y, 3) for y in bn([40.0, 50.0, 60.0])]) # [-1.104, -0.051, 1.003]
print([round(y, 3) for y in bn([50.0])]) # [-0.051]The estimates are near the true 50 and 100, not exact: with momentum 0.1 they mostly reflect the last ten or so batches. In eval mode, the value 50 maps to -0.051 whether it arrives alone or with others. Keras uses the opposite convention (momentum=0.99 means keep 99 percent of the old value), which is a classic porting bug.
Step 4: Why small batches hurt batch norm
bn = BatchNorm1d()
print([round(y, 3) for y in bn([48.0, 52.0])]) # [-1.0, 1.0]
print([round(y, 3) for y in bn([10.0, 90.0])]) # [-1.0, 1.0]
rng = random.Random(1)
for size in (2, 32):
batch_means = [mean([rng.gauss(50.0, 10.0) for _ in range(size)]) for _ in range(1000)]
print(size, round(min(batch_means), 1), round(max(batch_means), 1))
# 2 26.7 71.0
# 32 44.7 56.0With two samples, training-mode batch norm maps every pair of distinct values to -1 and +1; a gap of 4 and a gap of 80 look identical. The batch mean also swings from 27 to 71 across batches of two, versus 45 to 56 for batches of 32. The network trains against noisy statistics, then meets running averages at eval time that describe a different distribution. With one sample and no spatial dimensions, PyTorch refuses outright: BatchNorm1d raises “Expected more than 1 value per channel when training”.
Step 5: Gamma, beta and RMSNorm
Forcing mean 0 and variance 1 could remove information the next layer needs, so both layers end with a learned affine transform, gamma * x_hat + beta, initialized to 1 and 0. Because gamma and beta can undo the normalization, the layer never loses expressive power. RMSNorm skips the mean subtraction and the beta, dividing only by the root mean square:
def rms_norm(row, gamma, eps=1e-6):
rms = math.sqrt(mean([x * x for x in row]) + eps)
return [g * x / rms for g, x in zip(gamma, row)]
def affine(row, gamma, beta):
return [g * x + b for x, g, b in zip(row, gamma, beta)]
row = X[2]
m, v = mean(row), var(row)
undo = affine(layer_norm([row])[0], [math.sqrt(v + 1e-5)] * 4, [m] * 4)
print([round(x, 3) for x in undo]) # [3.0, 0.0, 9.0, 0.0]
print([round(x, 3) for x in rms_norm(row, [1.0] * 4)]) # [0.632, 0.0, 1.897, 0.0]RMSNorm (Zhang and Sennrich, 2019) is cheaper and works as well as LayerNorm for large language models; LLaMA-family models use it, and PyTorch has shipped nn.RMSNorm since 2.4. A related detail: a linear or convolution layer feeding batch norm usually sets bias=False, because the mean subtraction cancels any constant bias and beta replaces it.
Worked scenario
An image classifier with batch norm scores 94 percent on the validation set during training. Deployed behind an API that batches whatever requests arrive within 20 ms, accuracy falls to around 70 percent and the same image sometimes gets a different label. The serving code loads the weights but never calls model.eval(), so every batch norm layer is still in training mode:
def predict_one(bn, x, batch_mates):
return round(bn([x] + batch_mates)[0], 3)
bn = BatchNorm1d()
bn.running_mean, bn.running_var = 50.0, 100.0 # what training learned
print(predict_one(bn, 55.0, [40.0, 45.0]), predict_one(bn, 55.0, [70.0, 80.0])) # 1.336 -1.298
print(round(bn.running_mean, 2)) # 51.53, inference traffic moved the statistics
bn.running_mean, bn.running_var = 50.0, 100.0
bn.training = False # what model.eval() does for this layer
print(predict_one(bn, 55.0, [40.0, 45.0]), predict_one(bn, 55.0, [70.0, 80.0])) # 0.5 0.5The same input gets opposite activations depending on its neighbours, and production traffic overwrites the running statistics. The fix in PyTorch 2.x (illustrative):
model.eval() # batch norm uses running stats, dropout is disabled
with torch.inference_mode(): # no autograd bookkeeping
logits = model(images)The fine-tuning variant of this bug is a pretrained detector fine-tuned on one GPU with batch size 2. The usual fix is to freeze the batch norm layers so they keep their pretrained statistics, re-applied after every model.train() call because train() resets all submodules (PyTorch 2.x, illustrative):
model.train()
for m in model.modules():
if isinstance(m, (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d)):
m.eval() # use and keep the pretrained running stats
for p in m.parameters():
p.requires_grad_(False) # optionally freeze gamma and beta as wellAlternatives are GroupNorm (normalizes groups of channels within each sample) or, with several GPUs, SyncBatchNorm.
Common mistake
- “Gradient accumulation fixes small-batch batch norm.” It enlarges the effective batch for the gradient, but each forward pass still normalizes with statistics from its own micro-batch.
- “Layer norm normalizes over the batch dimension for each feature.” That is batch norm. Layer norm uses one sample’s features.
- “Eval mode only matters for dropout.” It switches batch norm to running statistics too, and calling
torch.no_grad()does not do that.
Verify the behavior
Append these assertions to the snippets above and run the file with plain python3:
def test_batch_norm_normalizes_each_feature():
for col in zip(*batch_norm(X)):
assert abs(mean(col)) < 1e-9 and abs(var(col) - 1) < 1e-4
def test_layer_norm_normalizes_each_sample():
for row in layer_norm(X):
assert abs(mean(row)) < 1e-9 and abs(var(row) - 1) < 1e-4
def test_layer_norm_ignores_batch_mates_and_scale():
assert layer_norm([X[1]]) == layer_norm(X)[1:2]
scaled = layer_norm([[10 * x for x in X[0]]])[0]
assert all(abs(a - b) < 1e-5 for a, b in zip(scaled, layer_norm([X[0]])[0]))
def test_eval_mode_is_independent_of_the_batch():
bn = BatchNorm1d(); bn.training = False
assert bn([5.0, 1.0])[0] == bn([5.0, 100.0])[0]
for test in (test_batch_norm_normalizes_each_feature, test_layer_norm_normalizes_each_sample,
test_layer_norm_ignores_batch_mates_and_scale, test_eval_mode_is_independent_of_the_batch):
test()
print("ok")In PyTorch, run one input inside batches with different companions after model.train() and again after model.eval(), and compare outputs with torch.allclose.
Follow-up questions
Why do transformers use LayerNorm rather than batch norm? Batched sequences have different lengths and padding, so batch statistics mix real tokens with padding and change with batch composition, and generation often runs one sequence at a time. Layer norm works per token, so training and inference compute exactly the same function.
What is the difference between pre-LN and post-LN? The original Transformer applied LayerNorm(x + sublayer(x)) after each residual addition (post-LN). Most modern models compute x + sublayer(LayerNorm(x)) (pre-LN), which keeps an unnormalized residual path. Xiong et al. (2020) showed pre-LN gradients are better behaved at initialization, so deep stacks train more stably and depend less on learning-rate warmup. Pre-LN models add one final LayerNorm after the last block.
Can you fuse batch norm into the preceding convolution? Yes, at inference. With frozen running statistics, batch norm is a fixed per-channel scale and shift, which folds into the convolution’s weights and bias.
Interview exercise
A ResNet-based detector was pretrained with 8 images per GPU. You fine-tune it on a single GPU with 2 images per step and gradient accumulation over 32 steps, reasoning that the effective batch of 64 matches the original. Training loss is noisy, and validation accuracy after fine-tuning is worse than the pretrained model’s. What is going on, and what do you change?
Answer and reasoning
Gradient accumulation makes the gradient look like a batch of 64, but batch norm never sees 64 images: each forward pass normalizes with the mean and variance of 2 images. From Step 4, two samples per batch reduce each feature to roughly plus or minus one, and the batch statistics swing widely. The running averages are also being rewritten from these noisy micro-batches, so eval mode inherits the damage. The original run normalized over 8 images per GPU (unless it used SyncBatchNorm), so that is the scale the statistics were tuned for. The practical fix is to freeze batch norm (eval mode plus frozen gamma and beta, re-applied after every model.train()), keeping the pretrained statistics. If the new data differs a lot from the pretraining data, replace batch norm with GroupNorm and fine-tune longer, or get more images per forward pass with mixed precision or smaller crops.