Modern AI Engineering

Lesson 3.6 · 23 min

Batch Norm vs Layer Norm: When to Use Each

Two layers do almost the same arithmetic, subtract a mean and divide by a standard deviation, so why does one rule image networks and the other rule every Transformer?

In short: Normalization layers rescale activations so they have a mean of about 0 and a standard deviation of about 1, then apply a learned scale and shift. Batch Normalization computes those statistics for each feature across the examples in a mini-batch; Layer Normalization computes them for each example across its features. That single difference in direction decides where each one works best.

What is normalization?

To normalize (more precisely, standardize) a set of numbers means to shift and rescale them so their mean (average) is 0 and their standard deviation (typical distance from the mean) is 1. You compute the mean μ and the variance σ² (the average squared distance from the mean), then transform each value:

Example: the values [2, 4, 6] have mean 4 and variance ((−2)² + 0² + 2²)/3 = 8/3, so standard deviation ≈ 1.63. Normalized: [−1.22, 0, 1.22].

The learned γ and β matter. They let the network undo the normalization if that is what works best: if it learns γ = σ and β = μ, it gets the original values back. So normalization never removes capacity; it just gives each layer a well-behaved starting point.

Think of it like grading on a curve Two teachers mark the same class: one gives scores out of 10, the other out of 1,000, and one marks harshly. Raw scores are hard to compare or combine. Converting each to "how many standard deviations above the average" puts them on the same scale. Normalization does that for the numbers flowing between layers.

Why do we need normalization?

In a deep network, each layer's output is the next layer's input. As weights change during training, the range of these activations can drift: some grow large, some shrink towards zero. That causes several problems:

  • Unstable gradients. Very large or very small activations lead to exploding or vanishing gradients, especially through saturating activations like sigmoid and tanh.
  • Sensitivity to the learning rate. If scales vary wildly, a learning rate that suits one layer is too big for another. Normalized networks tolerate larger learning rates and train faster.
  • Sensitivity to initialisation. Without normalization, deep networks only train if the initial weights are chosen very carefully.
  • A moving target. Each layer must keep adapting to the changing distribution of its inputs. The BatchNorm paper called this internal covariate shift. Later research questioned whether that is the real reason BatchNorm helps and argued it mainly makes the loss landscape smoother. Either way, the practical benefit is well established.

Normalization layers address all of these by resetting activations to a standard scale at chosen points in the network.

What is Batch Normalization?

Batch Normalization (BatchNorm, BN) was introduced by Sergey Ioffe and Christian Szegedy in 2015. For each feature (each neuron output, or each channel in a convolutional network) it computes the mean and variance across all the examples in the current mini-batch, and normalizes that feature with them.

Picture the activations as a table: rows are examples in the batch, columns are features. BatchNorm normalizes down each column. Feature 2 of example 1 is compared with feature 2 of the other examples in the batch.

BatchNorm, step by step

  1. Collect a column: For feature j, gather its value from every example in the mini-batch (for a CNN, also from every spatial position).
  2. Batch statistics: Compute that column's mean μⱼ and variance σⱼ².
  3. Normalize: x̂ = (x − μⱼ) / √(σⱼ² + ε) for every value in the column.
  4. Scale and shift: y = γⱼ · x̂ + βⱼ with the feature's learned parameters.
  5. Track running statistics: Update an exponential moving average of μⱼ and σⱼ² (PyTorch uses momentum 0.1 by default). These are saved with the model.
  6. At inference: Use the saved running mean and variance instead of batch statistics, so a single example's output does not depend on whatever else is in its batch.

That last step is BatchNorm's biggest quirk: it behaves differently in training and inference. During training an example's output depends on the other examples in its mini-batch. This creates a mild regularising noise (useful), but also problems when batches are small or when training and test data differ.

What is Layer Normalization?

Layer Normalization (LayerNorm, LN) was proposed by Jimmy Ba, Jamie Kiros and Geoffrey Hinton in 2016. It computes the mean and variance across the features of a single example, and normalizes that example with its own statistics. In our table, it normalizes across each row.

In a Transformer, every token has a hidden vector of, say, 768 numbers. LayerNorm takes those 768 numbers for one token, computes their mean and variance, and normalizes them. Each token in each sequence is handled independently. The learned γ and β still have one value per feature (768 each).

Because the statistics come from the example itself, LayerNorm does exactly the same thing in training and inference, needs no running averages, and works with any batch size, including 1. It also handles sequences of different lengths naturally, since each token is normalized on its own.

Pause and think: A Transformer processes a batch of 8 sequences × 128 tokens with hidden size 512. How many separate mean/variance pairs does LayerNorm compute in one layer?

One per token: 8 × 128 = 1,024 pairs, each over 512 numbers. BatchNorm on the same tensor (treating hidden units as features) would instead compute 512 pairs, each over the 1,024 token positions.

Both, side by side in code

The two functions below differ in a single argument: axis=0 (down the columns, BatchNorm) versus axis=1 (across the rows, LayerNorm). The last two lines show what happens with a batch of one example.

bn_vs_ln.py

import numpy as np
# A mini-batch: 4 examples (rows) x 3 features (columns)
X = np.array([[1.0, 200.0, 0.5],
[2.0, 220.0, 0.1],
[3.0, 180.0, 0.9],
[4.0, 240.0, 0.3]])
eps = 1e-5
gamma, beta = np.ones(3), np.zeros(3)      # learnable scale and shift
def batch_norm(X):
mu = X.mean(axis=0)                    # one mean per FEATURE (down columns)
var = X.var(axis=0)
return gamma * (X - mu) / np.sqrt(var + eps) + beta
def layer_norm(X):
mu = X.mean(axis=1, keepdims=True)     # one mean per EXAMPLE (across a row)
var = X.var(axis=1, keepdims=True)
return gamma * (X - mu) / np.sqrt(var + eps) + beta
np.set_printoptions(precision=2, suppress=True)
B, L = batch_norm(X), layer_norm(X)
print("BatchNorm output:\n", B)
print("  column means", B.mean(axis=0).round(2) + 0, " column stds", B.std(axis=0))
print("LayerNorm output:\n", L)
print("  row means   ", L.mean(axis=1).round(2) + 0, " row stds   ", L.std(axis=1))
# Batch of ONE example: BatchNorm has nothing to compare against
one = X[:1]
print("BatchNorm, batch size 1:", batch_norm(one))
print("LayerNorm, batch size 1:", layer_norm(one))

Output:

BatchNorm output:
[[-1.34 -0.45  0.17]
[-0.45  0.45 -1.18]
[ 0.45 -1.34  1.52]
[ 1.34  1.34 -0.51]]
column means [0. 0. 0.]  column stds [1. 1. 1.]
LayerNorm output:
[[-0.7   1.41 -0.71]
[-0.7   1.41 -0.72]
[-0.69  1.41 -0.72]
[-0.69  1.41 -0.72]]
row means    [0. 0. 0. 0.]  row stds    [1. 1. 1. 1.]
BatchNorm, batch size 1: [[0. 0. 0.]]
LayerNorm, batch size 1: [[-0.7   1.41 -0.71]]

Look at the LayerNorm rows. They are almost identical, because the big feature (≈ 200) dominates every row's mean and variance. This is a useful warning: LayerNorm assumes the features of one example are comparable, which is true for hidden units inside a network, but not for raw input columns measured in different units. BatchNorm, by normalizing each feature separately, handled the mixed scales well.

Batch Normalization vs Layer Normalization

When to use which one?

Real-world use ResNet and many classic image models place BatchNorm after almost every convolution. The original Transformer, BERT and GPT-2 use LayerNorm. GPT-2 moved LayerNorm to the start of each sub-layer ("pre-norm") rather than after the residual addition ("post-norm"), which makes deep Transformers more stable to train; pre-norm is now the common choice. Many recent LLMs replaced LayerNorm with RMSNorm, covered next.

Common mistakes Running inference with a BatchNorm model still in training mode: predictions then depend on the batch and can change wildly for a single example. Training BatchNorm with tiny batches (e.g. 2) and wondering why results are noisy. Fine-tuning with a very different data distribution and forgetting that BatchNorm's running statistics also need updating or freezing deliberately. And applying LayerNorm to raw features measured in different units.

Pause and think: You fine-tune an image model that uses BatchNorm, but your GPU only fits a batch size of 2. What two options could you consider?

Freeze the BatchNorm layers (keep their pretrained running statistics and parameters fixed, using evaluation-mode behaviour), or replace them with a batch-independent alternative such as GroupNorm. Gradient accumulation alone does not help, because BatchNorm statistics are still computed over each tiny batch.

Worked example, step by step

We said BatchNorm keeps an exponential moving average of the mean and variance for use at inference. Let us follow that average with real numbers, because it is the source of the most confusing BatchNorm bugs.

Take one feature whose true mean is 10. The running mean starts at 0 (the default) and the momentum is 0.1. Each training batch updates it with running ← 0.9 · running + 0.1 · batch_mean. To keep it simple, suppose every batch mean is exactly 10.

The running mean warming up

  1. Batch 1: 0.9 · 0 + 0.1 · 10 = 1.0. After one batch the stored mean is still 9 away from the truth.
  2. Batch 2: 0.9 · 1.0 + 0.1 · 10 = 1.9.
  3. Batch 3: 0.9 · 1.9 + 0.1 · 10 = 2.71. Each update closes one tenth of the remaining gap.
  4. The pattern: After n batches the running mean is 10 · (1 − 0.9ⁿ). The gap shrinks by a factor of 0.9 per batch.
  5. How long until it is right?: After 10 batches: 6.51. After 22 batches: 9.02. After 44 batches: 9.90. It takes dozens of batches before inference statistics can be trusted.

Now the consequence. In training mode the value 12 is normalized with the batch mean 10, so it comes out a little above zero. In evaluation mode after only three batches, the same 12 is normalized with the stored mean 2.71 and comes out far above zero. The layers after it have never seen such numbers.

Symptoms that point to running statistics
What we seeLikely causeWhat to check
Good training loss, terrible validation loss in the first few hundred stepsRunning statistics have not caught up yetValidate again later; compare with a run in training mode
A fine-tuned model is worse in evaluation mode than in training modeStored statistics still describe the old dataLet them update on the new data, or freeze the layers on purpose
Results change with the batch size at inferenceThe model is still in training modeSwitch to evaluation mode

LayerNorm has none of these problems. It stores no statistics, so there is nothing to warm up and nothing to go stale.

Practice: try it yourself

We build a tiny BatchNorm for a single feature, with a training mode and an evaluation mode. Then we watch two things: how the same value gets different outputs depending on its batch mates, and how the running statistics slowly become usable.

practice_batchnorm_modes.py

import numpy as np
class BatchNorm1Feature:
"""BatchNorm for a single feature, without the learned scale and shift."""
def __init__(self, momentum=0.1, eps=1e-5):
self.run_mean, self.run_var = 0.0, 1.0    # starting values
self.momentum, self.eps = momentum, eps
self.training = True
def __call__(self, x):
if self.training:                         # use this batch's statistics
mean, var = x.mean(), x.var()
m = self.momentum                     # and update the running ones
self.run_mean = (1 - m) * self.run_mean + m * mean
self.run_var = (1 - m) * self.run_var + m * var
else:                                     # use the stored statistics
mean, var = self.run_mean, self.run_var
return (x - mean) / np.sqrt(var + self.eps)
rng = np.random.default_rng(0)
bn = BatchNorm1Feature()
# 1) Training mode: the same value 12.0 with two different sets of batch mates
print("train, mates 8 and 10 :", bn(np.array([12.0, 8.0, 10.0])).round(2))
print("train, mates 14 and 16:", bn(np.array([12.0, 14.0, 16.0])).round(2))
# 2) Feed batches drawn around mean 10, std 2, then test 12.0 in eval mode
bn = BatchNorm1Feature()
for step in range(1, 101):
bn.training = True
bn(rng.normal(10.0, 2.0, size=32))
if step in (1, 5, 20, 100):
bn.training = False
out = bn(np.array([12.0]))[0]
print(f"after {step:3d} batches: run_mean={bn.run_mean:5.2f} "
f"run_var={bn.run_var:4.2f}  eval(12.0)={out:5.2f}")

Output:

train, mates 8 and 10 : [ 1.22 -1.22  0.  ]
train, mates 14 and 16: [-1.22  0.    1.22]
after   1 batches: run_mean= 0.97 run_var=1.16  eval(12.0)=10.26
after   5 batches: run_mean= 4.10 run_var=2.07  eval(12.0)= 5.48
after  20 batches: run_mean= 8.75 run_var=3.50  eval(12.0)= 1.74
after 100 batches: run_mean= 9.99 run_var=4.15  eval(12.0)= 0.98

Now change it:

  • Create the second layer with BatchNorm1Feature(momentum=0.5). Predict how the after 5 batches line changes. Then think about the price: with a large momentum, what does one unusual batch do to the stored statistics?
  • Change size=32 to size=2. The true variance is 4. Predict whether run_var still ends near 4, and whether eval(12.0) ends above or below 1.
  • After the loop, evaluate bn(np.array([22.0])), a typical-plus-one-std value from new data centred on 20. Predict the output, and say what it tells us about using old running statistics on data that has shifted.

Pause and think: In training mode the value 12.0 came out as +1.22 with one set of batch mates and −1.22 with another. Is that a bug in our layer?

No, it is BatchNorm doing exactly what it is defined to do. In training mode the output says where a value sits relative to its batch: 12 is the top of [12, 8, 10] and the bottom of [12, 14, 16]. This is why an example's output depends on its mini-batch during training, and why inference must switch to fixed, stored statistics to give one stable answer per input.

Pause and think: After one batch, evaluation mode turns 12.0 into 10.26, although the correct normalized value is about 1. A teammate concludes that the model is broken. What would we tell them?

The model is fine; the stored statistics are not ready. After one batch the running mean is 0.97 instead of 10 and the running variance 1.16 instead of 4, so (12 − 0.97) / √1.16 is about 10. The running averages move only a tenth of the way per batch. Evaluate after more training steps, and the same input gives about 1.

Key takeaways

  • Normalization subtracts a mean, divides by a standard deviation, then applies a learned scale γ and shift β.
  • BatchNorm: statistics per feature across the mini-batch; uses running averages at inference.
  • LayerNorm: statistics per example across its features; identical in training and inference.
  • BatchNorm suits CNNs with reasonable batch sizes; LayerNorm suits Transformers, RNNs and small batches.
  • Always use evaluation mode for BatchNorm at inference, and never LayerNorm raw features of different units.

Key terms

  • Normalization: Rescaling a group of values to mean 0 and standard deviation 1, then applying a learned scale and shift.
  • Batch Normalization: Normalizing each feature using the mean and variance across the examples of a mini-batch.
  • Layer Normalization: Normalizing each example using the mean and variance across its own features.
  • Running statistics: Moving averages of BatchNorm's mean and variance collected during training and used at inference.
  • γ and β: Learned per-feature scale and shift applied after normalization.
  • GroupNorm: A batch-independent variant that normalizes groups of channels within each example.

← 3.5 Dropout: Controlled Forgetting as Regularization · 3.7 RMSNorm: Simpler Normalization for Transformers →