Lesson 46: Batch Normalization

USAAIO Lesson 46, from Phase 2. It covers the batch-normalization forward pass, with mu, sigma^2, x_hat, and the gamma and beta parameters, then the full backward pass giving dL/dgamma, dL/dbeta, and dL/dx, and the running mean and variance used at inference. It then covers LayerNorm, which normalizes along the feature dimension for transformers, and GroupNorm for small batches. You implement BatchNorm1d from scratch and verify it against nn.BatchNorm1d. The lesson runs to 31 slides.

Subject: Machine Learning · 60 slides · code lesson

Open the interactive version of this deck · Homework for this lesson

What this lesson covers

The lesson, slide by slide

1. Batch Normalization: Forward, Backward & Variants

Title

USAAIO · Lesson 46 · Phase 2

Derive every gradient from scratch, implement running statistics for inference, and contrast BatchNorm with LayerNorm (transformers) and GroupNorm (small batches).

2. By the end of this lesson you can

Objectives

  1. Compute the BatchNorm forward pass (mu, sigma^2, x_hat, gamma, beta) and verify against nn.BatchNorm1d
  2. Derive all three backward-pass gradients: dL/dgamma, dL/dbeta, dL/dx
  3. Explain the running mean/variance momentum update and why inference uses it instead of batch stats
  4. Distinguish LayerNorm (feature-dim) from BatchNorm (batch-dim) and state when each is preferred
  5. Implement BatchNorm1d from scratch including training and eval modes

3. What survived from Learning Curves, Bias-Variance Diagnosis, and Double Descent?

Warm-up

Discussion prompt

Before we open Lesson 46: Batch Normalization: without looking back, what was the main idea of Learning Curves, Bias-Variance Diagnosis, and Double Descent, and what could you do by the end of it that you could not do before?

Hint: One sentence for the idea, one for the skill. If the second one is blank, that is the part to revisit.

Answer:

plot train and val loss vs epochs and vs training-set size; diagnose underfitting (high bias) and overfitting (high variance) from curve shapes; apply cures — dropout, early stopping with patience, data augmentation, capacity changes; understand the double-descent phenomenon for overparameterized models.

4. The forward pass

Section

Part 1 of 4

5. Why normalize activations?

Concept

Deep layers see inputs whose distribution shifts as earlier weights update — internal covariate shift. Each layer must adapt to a moving target.

BatchNorm forces each feature's activations to have zero mean and unit variance within a mini-batch, then lets the network re-scale via learned gamma and beta. The layer sees a stable distribution; gradients flow much better.

Empirical payoff (Lesson 32 gradient-flow context): a 10-layer MLP without BatchNorm stays stuck at loss = 0.693 (random); with BatchNorm it trains to 0.168 within 100 epochs on the same data.

6. Break it if you can: Why normalize activations?

Counterexample

Discussion prompt

Empirical payoff (Lesson 32 gradient-flow context): a 10-layer MLP without BatchNorm stays stuck at loss = 0.693 (random); with BatchNorm it trains to 0.168 within 100 epochs on the same data.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

7. The four-step forward pass

Concept

For a mini-batch x of shape (N, D) — N samples, D features — BatchNorm normalizes each feature independently over the N samples.

\[ \mu_j = \frac{1}{N}\sum_{i=1}^{N} x_{ij} \qquad \sigma^2_j = \frac{1}{N}\sum_{i=1}^{N}(x_{ij}-\mu_j)^2 \]

\[ \hat{x}_{ij} = \frac{x_{ij}-\mu_j}{\sqrt{\sigma^2_j+\varepsilon}} \qquad y_{ij} = \gamma_j\,\hat{x}_{ij} + \beta_j \]

gamma and beta are learned parameters (one per feature). With gamma=1, beta=0 the output is the raw standardized x_hat. Training can recover any other scale/shift.

8. By analogy: The four-step forward pass

Analogy

Discussion prompt

Explain The four-step forward pass by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

For a mini-batch x of shape (N, D) — N samples, D features — BatchNorm normalizes each feature independently over the N samples.

9. What has to happen first: Forward pass trace: 3 samples, 1 feature

Ranking

Put in order

Put the moves of Forward pass trace: 3 samples, 1 feature into the order they have to happen.

  1. Step 1: mu = (1+4+7)/3 = 4.0
  2. Step 2: var = ((1-4)^2 + (4-4)^2 + (7-4)^2)/3 = 6.0, sqrt(var+eps) = 2.4495
  3. Step 3: x_hat = [(-3)/2.4495, 0/2.4495, 3/2.4495] = [-1.2247, 0.0, 1.2247]

Why: These are the moves of the worked example in the order it makes them, and each one is set up by the one before it. Population variance (unbiased=False) over the batch — each sample's squared deviation from mu.

10. Forward pass trace: 3 samples, 1 feature

Worked example

Input batch x = [1.0, 4.0, 7.0] (N=3), gamma=2, beta=-1. Compute the BatchNorm output step-by-step.

import torch, torch.nn as nn
x = torch.tensor([[1.0], [4.0], [7.0]])   # shape (3, 1)
bn = nn.BatchNorm1d(1, eps=1e-5)
bn.weight.data.fill_(2.0)   # gamma = 2
bn.bias.data.fill_(-1.0)    # beta  = -1
print(bn(x).detach().squeeze().tolist())

Step 1: mu = (1+4+7)/3 = 4.0

Why: Mean over the batch dimension.

Step 2: var = ((1-4)^2 + (4-4)^2 + (7-4)^2)/3 = 6.0, sqrt(var+eps) = 2.4495

Why: Population variance (unbiased=False) over the batch — each sample's squared deviation from mu.

Step 3: x_hat = [(-3)/2.4495, 0/2.4495, 3/2.4495] = [-1.2247, 0.0, 1.2247]

Why: Standardize: subtract mu, divide by sqrt(var+eps).

samplexx - mux_haty = 2*x_hat - 1
11.0-3.0-1.2247-3.4495
24.00.00.0000-1.0000
37.0+3.0+1.2247+1.4495

11. Fill in: x for Forward pass trace: 3 samples, 1 feature

Comparison

Comparison matrix

From Forward pass trace: 3 samples, 1 feature: refill the x column from what you know. The rest of the table is as it appeared.

samplexx - mux_haty = 2*x_hat - 1
11.0-3.0-1.2247-3.4495
24.00.00.0000-1.0000
37.0+3.0+1.2247+1.4495

12. The backward pass

Section

Part 2 of 4

13. Gradient flow through BatchNorm

Concept

BatchNorm sits inside the graph like any other op. The backward pass chains through y = gamma*x_hat + beta, then x_hat = (x-mu)/sigma, then sigma and mu are both functions of x.

\[ \frac{\partial L}{\partial \gamma_j} = \sum_{i=1}^{N} \frac{\partial L}{\partial y_{ij}}\,\hat{x}_{ij} \qquad \frac{\partial L}{\partial \beta_j} = \sum_{i=1}^{N} \frac{\partial L}{\partial y_{ij}} \]

These two are immediate chain-rule applications. The hard gradient is dL/dx — it must account for the fact that changing any single x_ij also shifts the batch mean and variance, which affects all other samples.

14. Teach it back: Gradient flow through BatchNorm

Explain it

Discussion prompt

Explain Gradient flow through BatchNorm to a student a year behind you. No notation, no jargon they have not met — and it still has to be true.

Hint: If your explanation needs a symbol they have never seen, you are describing the notation rather than the idea.

Answer:

BatchNorm sits inside the graph like any other op. The backward pass chains through y = gamma*x_hat + beta, then x_hat = (x-mu)/sigma, then sigma and mu are both functions of x.

15. dL/dx: the full chain

Concept

Let g = dL/dy * gamma (upstream signal scaled by gamma). Then x_hat's gradient is dx_hat = dL/dy * gamma.

\[ \frac{\partial L}{\partial x_{ij}} = \frac{1}{N}\,\frac{1}{\sqrt{\sigma^2_j+\varepsilon}}\Bigl(N\,g_{ij} - \sum_{k}g_{kj} - \hat{x}_{ij}\sum_{k}g_{kj}\hat{x}_{kj}\Bigr) \]

The two subtracted terms cancel the change that x_ij induces in mu and sigma. Without them, the gradient pretends mu and sigma are constants — which they are not during training.

16. What rests on this: dL/dx: the full chain

Socratic

Discussion prompt

Let g = dL/dy * gamma (upstream signal scaled by gamma). Then x_hat's gradient is dx_hat = dL/dy * gamma.

Suppose that were not true. What is the first thing in Lesson 46: Batch Normalization that would stop working?

Hint: Follow it one step downstream. The answer is whatever was quietly relying on it.

Answer:

The two subtracted terms cancel the change that x_ij induces in mu and sigma. Without them, the gradient pretends mu and sigma are constants — which they are not during training.

17. Guess the shape of the answer: Verify dL/dgamma, dL/dbeta, dL/dx with…

Estimation

Predict first

Compute the three gradients manually and verify them against PyTorch autograd.

Commit before you compute: what does Verify dL/dgamma, dL/dbeta, dL/dx with autograd come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: autograd dL/dbeta = [-3.2438, 1.0020]

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Sum of dout over the batch dimension — the chain from y = ...

18. Verify dL/dgamma, dL/dbeta, dL/dx with autograd

Worked example

Compute the three gradients manually and verify them against PyTorch autograd.

import torch
torch.manual_seed(0)
N, D = 4, 2
x = torch.randn(N, D, requires_grad=True)
gamma = torch.tensor([2., 0.5], requires_grad=True)
beta  = torch.tensor([0.1, -0.3], requires_grad=True)
dout  = torch.randn(N, D)
eps = 1e-5
mu   = x.mean(0); var = x.var(0, unbiased=False)
x_hat = (x - mu) / (var + eps).sqrt()
y = gamma * x_hat + beta
y.backward(dout)
print('dL/dgamma:', gamma.grad.tolist())
print('dL/dbeta: ', beta.grad.tolist())

autograd dL/dgamma = [-0.2690, -1.3097]

Why: Sum of (dout * x_hat) over the batch dimension — the chain from y = gamma*x_hat.

autograd dL/dbeta = [-3.2438, 1.0020]

Why: Sum of dout over the batch dimension — the chain from y = ... + beta.

gradientautogradmanual formulamatch
dL/dgamma[0]-0.2690sum(dout*x_hat, axis=0)[0]True
dL/dbeta[0]-3.2438sum(dout, axis=0)[0]True
dL/dx[0,0]0.25481/N * 1/std * (Ng - sum_g - x_hatsum_gx)True

19. What each one costs: Verify dL/dgamma, dL/dbeta, dL/dx with autograd

Trade off

Comparison matrix

From Verify dL/dgamma, dL/dbeta, dL/dx with autograd: every row here is a choice with a cost. Fill the match column, then say which row you would actually pick and what you give up for it.

gradientautogradmanual formulamatch
dL/dgamma[0]-0.2690sum(dout*x_hat, axis=0)[0]True
dL/dbeta[0]-3.2438sum(dout, axis=0)[0]True
dL/dx[0,0]0.25481/N * 1/std * (Ng - sum_g - x_hatsum_gx)True

20. Something is wrong here: treating mu and sigma as constants in dL/dx

Anomaly

Predict first

A student writes this, and it looks reasonable:

x_hat = (x - mu) / sigma, so dx_hat/dx_i = 1/sigma. Therefore dL/dx_i = dL/dx_hat * gamma / sigma.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: This treats mu and sigma as fixed constants.

Use the full chain rule: x_hat depends on x both directly and through mu and sigma.

Why: This treats mu and sigma as fixed constants. But both are functions of the entire batch: changing x_i changes mu and sigma for every sample in the batch.

21. Trap: treating mu and sigma as constants in dL/dx

Trap

The trap

x_hat = (x - mu) / sigma, so dx_hat/dx_i = 1/sigma. Therefore dL/dx_i = dL/dx_hat * gamma / sigma.

dL/dx_i = g_i / sigma (forgetting mu and sigma depend on x)

Why: This treats mu and sigma as fixed constants. But both are functions of the entire batch: changing x_i changes mu and sigma for every sample in the batch.

The fix

Use the full chain rule: x_hat depends on x both directly and through mu and sigma.

dL/dx_i = (1/N) * (1/sigma) * (N*g_i - sum(g) - x_hat_i * sum(g * x_hat))

Why: The two subtracted terms cancel the batch-mean and batch-variance dependencies. Skipping them gives wrong gradients that fail to converge correctly.

22. Break it on purpose: treating mu and sigma as constants in dL/dx

Break the constraint

Discussion prompt

The rule this trap just fixed:

The two subtracted terms cancel the batch-mean and batch-variance dependencies. Skipping them gives wrong gradients that fail to converge correctly.

Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?

Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.

Answer:

This treats mu and sigma as fixed constants. But both are functions of the entire batch: changing x_i changes mu and sigma for every sample in the batch.

23. Running stats for inference

Section

Part 3 of 4

24. Why inference needs different statistics

Concept

At test time you often predict one sample at a time — there is no batch to compute mu and sigma from. Using a single sample's own mean/std would re-normalize it to zero, erasing information.

BatchNorm maintains running (exponential moving average) statistics during training. At inference (model.eval()), it uses the running stats instead of the batch stats.

\[ \mu_{\text{run}} \leftarrow (1-m)\,\mu_{\text{run}} + m\,\mu_{\text{batch}} \qquad \sigma^2_{\text{run}} \leftarrow (1-m)\,\sigma^2_{\text{run}} + m\,\sigma^2_{\text{batch}} \]

25. What rests on this: Why inference needs different statistics

Socratic

Discussion prompt

At test time you often predict one sample at a time — there is no batch to compute mu and sigma from. Using a single sample's own mean/std would re-normalize it to zero, erasing information.

Suppose that were not true. What is the first thing in Lesson 46: Batch Normalization that would stop working?

Hint: Follow it one step downstream. The answer is whatever was quietly relying on it.

Answer:

BatchNorm maintains running (exponential moving average) statistics during training. At inference (model.eval()), it uses the running stats instead of the batch stats.

26. Guess the shape of the answer: Running mean trace: 3 batches, momentum=0.1

Estimation

Predict first

Three consecutive batches each with batch mean 4.0, 5.0, 6.0. Running mean starts at 0 and converges toward the true mean.

Commit before you compute: what does Running mean trace: 3 batches, momentum=0.1 come out to? A rough magnitude and the right form is enough — the point is to have something concrete to be wrong about.

Correct: After batch 1: run_mean = 0.90.0 + 0.14.0 = 0.4000

Why: A prediction you can defend turns the computation into a check rather than a leap of faith — and an answer that contradicts it is caught on the spot. Exponential smoothing: the running mean starts near 0 and creeps toward the data distribution.

27. Running mean trace: 3 batches, momentum=0.1

Worked example

Three consecutive batches each with batch mean 4.0, 5.0, 6.0. Running mean starts at 0 and converges toward the true mean.

import torch, torch.nn as nn
bn = nn.BatchNorm1d(1, eps=1e-5, momentum=0.1)
batches = [
    torch.tensor([[1.0],[4.0],[7.0]]),   # mu=4.0
    torch.tensor([[2.0],[5.0],[8.0]]),   # mu=5.0
    torch.tensor([[3.0],[6.0],[9.0]]),   # mu=6.0
]
for b in batches:
    bn(b)
    print(f'run_mean={bn.running_mean.item():.4f}')

After batch 1: run_mean = 0.90.0 + 0.14.0 = 0.4000

Why: Exponential smoothing: the running mean starts near 0 and creeps toward the data distribution.

batchbatch_murun_mean after
14.00000.4000
25.00000.8600
36.00001.3740

28. Watch it run: Running mean trace: 3 batches, momentum=0.1

Pattern

Step through it

Step through Running mean trace: 3 batches, momentum=0.1 one row at a time. What is driving the change, and what would the row after the last one be?

  1. Step 1: batch is 1
  2. Step 2: batch is 2
  3. Step 3: batch is 3

29. Train mode vs eval mode

Concept

modenormalizes usingwhen usedcall
traincurrent batch mu, sigma^2parameter updatesmodel.train()
evalrunning mu, sigma^2inference / single samplemodel.eval()

Forgetting model.eval() before inference is a common bug: the model will compute per-sample statistics (degenerate for N=1) and produce wrong outputs.

30. Where does each piece belong: Lesson 46: Batch Normalization

Sorting

Sort into buckets

These are the pieces of Lesson 46: Batch Normalization, out of order. Put each one back under the part of the lesson it belongs to.

The forward pass
Why normalize activations?; The four-step forward pass; Forward pass trace: 3 samples, 1 feature
The backward pass
Gradient flow through BatchNorm; dL/dx: the full chain; Verify dL/dgamma, dL/dbeta, dL/dx with autograd
Running stats for inference
Why inference needs different statistics; Running mean trace: 3 batches, momentum=0.1; Train mode vs eval mode
s1
The forward pass is where Lesson 46: Batch Normalization puts Why normalize activations?, The four-step forward pass, Forward pass trace: 3 samples, 1 feature. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s2
The backward pass is where Lesson 46: Batch Normalization puts Gradient flow through BatchNorm, dL/dx: the full chain, Verify dL/dgamma, dL/dbeta, dL/dx with autograd. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.
s3
Running stats for inference is where Lesson 46: Batch Normalization puts Why inference needs different statistics, Running mean trace: 3 batches, momentum=0.1, Train mode vs eval mode. Knowing which part of the lesson a problem belongs to is most of knowing which method to reach for.

31. Something is wrong here: forgetting model.eval() at inference

Anomaly

Predict first

A student writes this, and it looks reasonable:

After training, immediately run inference with a single sample — the model is still in train mode.

It is wrong. Say what breaks — and say it before you turn the page.

Correct: In train mode with N=1, BatchNorm normalizes this single sample against itself: mu=x, var=0.

Call model.eval() before every inference block.

Why: In train mode with N=1, BatchNorm normalizes this single sample against itself: mu=x, var=0. Every feature becomes 0 (or NaN with eps). Outputs are meaningless.

32. Trap: forgetting model.eval() at inference

Trap

The trap

After training, immediately run inference with a single sample — the model is still in train mode.

x_test = torch.tensor([[2.0, 3.0]]); y = model(x_test)

Why: In train mode with N=1, BatchNorm normalizes this single sample against itself: mu=x, var=0. Every feature becomes 0 (or NaN with eps). Outputs are meaningless.

The fix

Call model.eval() before every inference block.

model.eval(); with torch.no_grad(): y = model(x_test)

Why: eval() switches BatchNorm to use the accumulated running statistics that represent the full training distribution. This gives correct, stable predictions for any batch size including N=1.

33. Which of these survive contact with Lesson 46: Batch Normalization?

Two truths and a lie

Sort into buckets

Some of these hold up and some are the exact mistakes this lesson is built to prevent. Sort them.

Holds up
For a mini-batch x of shape (N, D) — N samples, D features — BatchNorm normalizes each feature independently over the N samples.; Let g = dL/dy * gamma (upstream signal scaled by gamma). Then x_hat's gradient is dx_hat = dL/dy * gamma.; BatchNorm maintains running (exponential moving average) statistics during training. At inference (model.eval()), it uses the running stats instead of the batch stats.
Breaks
x_hat = (x - mu) / sigma, so dx_hat/dx_i = 1/sigma. Therefore dL/dx_i = dL/dx_hat * gamma / sigma.; After training, immediately run inference with a single sample — the model is still in train mode.
sound
These are stated as this lesson states them — each one survives the edge cases Lesson 46: Batch Normalization puts it through.
flawed
Each of these is lifted from a trap in this deck: reasonable-sounding, and wrong in a way that only shows up once you rely on it.

34. LayerNorm & GroupNorm

Section

Part 4 of 4

35. LayerNorm: normalize over the feature dimension

Concept

LayerNorm (used in transformers) normalizes each sample independently over its feature dimension — no batch dependency at all.

\[ \text{LN}(x_i) = \gamma \cdot \frac{x_i - \mu_i}{\sqrt{\sigma^2_i + \varepsilon}} + \beta \qquad \mu_i = \frac{1}{D}\sum_j x_{ij} \]

For x = [1, 4, 7, 10] (single sample, D=4): mu=5.5, var=11.25, LN output = [-1.3416, -0.4472, +0.4472, +1.3416]. Same formula as BatchNorm but over features, not batch.

36. Teach it back: LayerNorm: normalize over the feature dimension

Explain it

Discussion prompt

Explain LayerNorm: normalize over the feature dimension to a student a year behind you. No notation, no jargon they have not met — and it still has to be true.

Hint: If your explanation needs a symbol they have never seen, you are describing the notation rather than the idea.

Answer:

LayerNorm (used in transformers) normalizes each sample independently over its feature dimension — no batch dependency at all.

37. BatchNorm vs LayerNorm vs GroupNorm

Concept

normnormalizes overtypical usebatch-size dep?
BatchNormbatch dim (N) per featureCNNs, MLPsYes — breaks at N=1
LayerNormfeature dim (D) per sampleTransformers, NLPNo
GroupNormgroups of channels per sampleSmall-batch visionNo

A key exam fact: BatchNorm is undefined (var=0) for a single sample per feature during training. LayerNorm and GroupNorm do not have this problem — they are preferred whenever batch size is small or variable.

38. Fill in: typical use for BatchNorm vs LayerNorm vs GroupNorm

Comparison

Comparison matrix

From BatchNorm vs LayerNorm vs GroupNorm: refill the typical use column from what you know. The rest of the table is as it appeared.

normnormalizes overtypical usebatch-size dep?
BatchNormbatch dim (N) per featureCNNs, MLPsYes — breaks at N=1
LayerNormfeature dim (D) per sampleTransformers, NLPNo
GroupNormgroups of channels per sampleSmall-batch visionNo

39. GroupNorm: normalize within channel groups

Concept

GroupNorm splits the C channels into G groups of C/G each, then normalizes within each group for each sample independently. When G=C it is InstanceNorm; when G=1 it is LayerNorm.

For 4 channels with G=2 on input [[1,2,3,4],[5,6,7,8]] (one sample): group 1 (channels 0-1) normalizes over [1,2,5,6] and group 2 over [3,4,7,8]. The output has the same shape but is normalized within each group.

40. By analogy: GroupNorm: normalize within channel groups

Analogy

Discussion prompt

Explain GroupNorm: normalize within channel groups by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

GroupNorm splits the C channels into G groups of C/G each, then normalizes within each group for each sample independently. When G=C it is InstanceNorm; when G=1 it is LayerNorm.

41. Rebuild the recipe: The BatchNorm recipe

Ranking

Put in order

These are the steps of The BatchNorm recipe, scrambled. Put them back in order before the next slide shows you.

  1. Forward (train): mu = batch mean; var = batch var; x_hat = (x-mu)/sqrt(var+eps); y = gamma*x_hat + beta
  2. Running stats: update each batch: run_mu = (1-m)run_mu + mbatch_mu (same for var)
  3. Forward (eval): use run_mu, run_var instead of batch stats
  4. Backward: dL/dgamma = sum(dout*x_hat); dL/dbeta = sum(dout); dL/dx uses the full chain (3 terms)
  5. Choice of norm: BatchNorm for large-batch CNNs/MLPs; LayerNorm for transformers; GroupNorm for small-batch vision

Why: This is the order the recipe itself gives. Recalling the sequence without the slide in front of you is the difference between recognising the method and being able to run it — most of what goes wrong in practice is a step done out of turn.

42. The BatchNorm recipe

Pattern

  1. Forward (train): mu = batch mean; var = batch var; x_hat = (x-mu)/sqrt(var+eps); y = gamma*x_hat + beta
  2. Running stats: update each batch: run_mu = (1-m)run_mu + mbatch_mu (same for var)
  3. Forward (eval): use run_mu, run_var instead of batch stats
  4. Backward: dL/dgamma = sum(dout*x_hat); dL/dbeta = sum(dout); dL/dx uses the full chain (3 terms)
  5. Choice of norm: BatchNorm for large-batch CNNs/MLPs; LayerNorm for transformers; GroupNorm for small-batch vision

43. Where does it stop working: The BatchNorm recipe

Edge cases

Discussion prompt

The BatchNorm recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".

Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.

Answer:

  1. Forward (train): mu = batch mean; var = batch var; x_hat = (x-mu)/sqrt(var+eps); y = gamma*x_hat + beta
  2. Running stats: update each batch: run_mu = (1-m)run_mu + mbatch_mu (same for var)
  3. Forward (eval): use run_mu, run_var instead of batch stats
  4. Backward: dL/dgamma = sum(dout*x_hat); dL/dbeta = sum(dout); dL/dx uses the full chain (3 terms)
  5. Choice of norm: BatchNorm for large-batch CNNs/MLPs; LayerNorm for transformers; GroupNorm for small-batch vision

44. Rule out three: Check yourself — forward pass

Elimination

Eliminate the wrong options

A BatchNorm layer sees batch x = [2.0, 6.0] (N=2, D=1, eps=0, gamma=1, beta=0). What is the output for the first sample (x=2.0)?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. -1.0
  • B. 0.0
  • C. -0.5
  • D. -2.0

Survives elimination: A

Why: mu=(2+6)/2=4, var=((2-4)^2+(6-4)^2)/2=4, std=2. x_hat_1=(2-4)/2=-1.0. With gamma=1, beta=0, output=-1.0.

45. Check yourself — forward pass

Check

Trace through a tiny example.

Check your understanding

A BatchNorm layer sees batch x = [2.0, 6.0] (N=2, D=1, eps=0, gamma=1, beta=0). What is the output for the first sample (x=2.0)?

  • A. -1.0 (correct)
  • B. 0.0
  • C. -0.5
  • D. -2.0

Answer: A

Why: mu=(2+6)/2=4, var=((2-4)^2+(6-4)^2)/2=4, std=2. x_hat_1=(2-4)/2=-1.0. With gamma=1, beta=0, output=-1.0.

Why B tempts people
x=2 is below the mean (4), so x_hat is negative, not 0. x=mu would give 0.
Why C tempts people
Dividing by variance (4) instead of std (2): (2-4)/4=-0.5. BatchNorm divides by the standard deviation, not the variance.
Why D tempts people
Subtracting 2*std instead of dividing by std: (2-4)-2=-4, or forgetting to divide at all. The denominator is std=2, giving -2/2=-1, not -2.

46. Answer it before you see the options: Check yourself — backward pass

Prediction

Predict first

Which formula correctly gives dL/dbeta in BatchNorm?

Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.

Correct: dL/dbeta = sum(dL/dy, axis=batch)

Why: y = gamma*x_hat + beta. By the chain rule, dy/dbeta = 1 for each sample, so dL/dbeta = sum over the batch of dL/dy * 1 = sum(dL/dy, axis=batch). This is analogous to the bias gradient in a linear layer.

47. Check yourself — backward pass

Check

Which gradient is easiest and which is hardest?

Check your understanding

Which formula correctly gives dL/dbeta in BatchNorm?

  • A. dL/dbeta = sum(dL/dy, axis=batch) (correct)
  • B. dL/dbeta = sum(dL/dy * x_hat, axis=batch)
  • C. dL/dbeta = mean(dL/dy, axis=batch)
  • D. dL/dbeta = dL/dy * gamma

Answer: A

Why: y = gamma*x_hat + beta. By the chain rule, dy/dbeta = 1 for each sample, so dL/dbeta = sum over the batch of dL/dy * 1 = sum(dL/dy, axis=batch). This is analogous to the bias gradient in a linear layer.

Why B tempts people
That is dL/dgamma, not dL/dbeta. For gamma, dy/dgamma = x_hat, giving the x_hat factor.
Why C tempts people
The gradient sums (not averages) over the batch — you accumulate how much each sample's prediction error wanted to shift beta.
Why D tempts people
dL/dy * gamma gives dL/dx_hat, the upstream gradient with respect to the normalized value, not beta.

48. Rule out three: Check yourself — train vs eval

Elimination

Eliminate the wrong options

You deploy a BatchNorm model and forget model.eval(). What happens when you predict a single sample?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. BatchNorm uses the single sample's own mean/std, making x_hat = 0 for every feature
  • B. BatchNorm silently switches to running stats automatically
  • C. BatchNorm skips normalization for N=1 batches
  • D. The model raises a RuntimeError and stops

Survives elimination: A

Why: In train mode with N=1, mu=x and var=0 for each feature. x_hat=(x-mu)/sqrt(0+eps) collapses to near-zero. The normalization destroys the signal — outputs are wrong, not an error.

49. Check yourself — train vs eval

Check

What breaks at inference?

Check your understanding

You deploy a BatchNorm model and forget model.eval(). What happens when you predict a single sample?

  • A. BatchNorm uses the single sample's own mean/std, making x_hat = 0 for every feature (correct)
  • B. BatchNorm silently switches to running stats automatically
  • C. BatchNorm skips normalization for N=1 batches
  • D. The model raises a RuntimeError and stops

Answer: A

Why: In train mode with N=1, mu=x and var=0 for each feature. x_hat=(x-mu)/sqrt(0+eps) collapses to near-zero. The normalization destroys the signal — outputs are wrong, not an error.

Why B tempts people
PyTorch does NOT auto-switch based on batch size. Only model.eval() triggers the use of running statistics.
Why C tempts people
BatchNorm does not have a batch-size check or bypass. It blindly applies the formula, producing degenerate outputs.
Why D tempts people
PyTorch does not raise an error for this — it silently produces incorrect results, which is why it is a dangerous silent bug.

50. Your turn: build BatchNorm1d

Section

Project

51. Project: BatchNorm1d from scratch

Concept

Implement BatchNorm1dScratch as an nn.Module with gamma, beta, running stats, and correct train/eval behavior. Verify every output against nn.BatchNorm1d.

#requirementtool
1forward pass (train mode): batch mu, var, x_hat, gamma*x_hat+betatorch.mean, torch.var
2running stats update (momentum=0.1)register_buffer, manual EMA
3forward pass (eval mode): use running statsself.training flag
4verify all outputs match nn.BatchNorm1dtorch.allclose

Build rules: use nn.Parameter for gamma and beta, register_buffer for running stats (they track state but are not optimized), and check self.training to branch between modes.

52. Break it if you can: Project: BatchNorm1d from scratch

Counterexample

Discussion prompt

Implement BatchNorm1dScratch as an nn.Module with gamma, beta, running stats, and correct train/eval behavior. Verify every output against nn.BatchNorm1d.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

Answer:

Build rules: use nn.Parameter for gamma and beta, register_buffer for running stats (they track state but are not optimized), and check self.training to branch between modes.

53. Milestone 1 — forward pass (train mode)

Worked example

Your turn: implement the forward method for training. Predict what x_hat looks like for a constant-step batch.

Hint: x.mean(dim=0) for batch mu; x.var(dim=0, unbiased=False) for batch var; normalize with (x - mu) / (var + eps).sqrt().

import torch, torch.nn as nn
class BN1d(nn.Module):
    def __init__(self, D, eps=1e-5, momentum=0.1):
        super().__init__()
        self.gamma = nn.Parameter(torch.ones(D))
        self.beta  = nn.Parameter(torch.zeros(D))
        self.eps, self.momentum = eps, momentum
        self.register_buffer('running_mean', torch.zeros(D))
        self.register_buffer('running_var',  torch.ones(D))
    def forward(self, x):
        if self.training:
            mu  = x.mean(0); var = x.var(0, unbiased=False)
            self.running_mean = (1-self.momentum)*self.running_mean + self.momentum*mu.detach()
            self.running_var  = (1-self.momentum)*self.running_var  + self.momentum*var.detach()
        else:
            mu, var = self.running_mean, self.running_var
        return self.gamma * (x - mu) / (var + self.eps).sqrt() + self.beta
checkresult
output shape matches inputTrue
x_hat mean per feature approx 0True
x_hat std per feature approx 1True
matches nn.BatchNorm1d outputTrue

54. Milestone 2 — verify against nn.BatchNorm1d

Worked example

Your turn: copy gamma/beta into your module, run the same input through both, and confirm allclose.

Hint: copy with scratch.gamma.data.copy_(ref.weight.data) — BatchNorm1d stores gamma in .weight and beta in .bias.

torch.manual_seed(42)
x = torch.randn(8, 4)
ref    = nn.BatchNorm1d(4, eps=1e-5, momentum=0.1)
scratch = BN1d(4, eps=1e-5, momentum=0.1)
scratch.gamma.data.copy_(ref.weight.data)
scratch.beta.data.copy_(ref.bias.data)
y_ref  = ref(x)
y_mine = scratch(x)
print('output match:', torch.allclose(y_ref, y_mine, atol=1e-5))
print('run_mean match:', torch.allclose(ref.running_mean, scratch.running_mean, atol=1e-5))
comparisonresult
output match (atol=1e-5)True
running_mean matchTrue
scratch.training flagTrue (default)

55. What each one costs: Milestone 2 — verify against nn.BatchNorm1d

Trade off

Comparison matrix

From Milestone 2 — verify against nn.BatchNorm1d: every row here is a choice with a cost. Fill the result column, then say which row you would actually pick and what you give up for it.

comparisonresult
output match (atol=1e-5)True
running_mean matchTrue
scratch.training flagTrue (default)

56. The full program

Concept

import torch, torch.nn as nn
class BN1d(nn.Module):
    def __init__(self, D, eps=1e-5, momentum=0.1):
        super().__init__()
        self.gamma = nn.Parameter(torch.ones(D))
        self.beta  = nn.Parameter(torch.zeros(D))
        self.eps, self.momentum = eps, momentum
        self.register_buffer('running_mean', torch.zeros(D))
        self.register_buffer('running_var',  torch.ones(D))
    def forward(self, x):
        if self.training:
            mu  = x.mean(0); var = x.var(0, unbiased=False)
            self.running_mean = (1-self.momentum)*self.running_mean + self.momentum*mu.detach()
            self.running_var  = (1-self.momentum)*self.running_var  + self.momentum*var.detach()
        else:
            mu, var = self.running_mean, self.running_var
        return self.gamma*(x-mu)/(var+self.eps).sqrt() + self.beta
torch.manual_seed(42); x = torch.randn(8, 4)
ref = nn.BatchNorm1d(4); s = BN1d(4)
s.gamma.data.copy_(ref.weight.data); s.beta.data.copy_(ref.bias.data)
print('match:', torch.allclose(ref(x), s(x), atol=1e-5))
componentimplementation detail
gamma / betann.Parameter — learned by the optimizer
running_mean/varregister_buffer — saved in state_dict, not optimized
train vs evalself.training flag, toggled by .train() / .eval()

If match: True — you have implemented the same BatchNorm that every ResNet, BERT encoder, and GPT block uses.

57. Fill in: implementation detail for The full program

Comparison

Comparison matrix

From The full program: refill the implementation detail column from what you know. The rest of the table is as it appeared.

componentimplementation detail
gamma / betann.Parameter — learned by the optimizer
running_mean/varregister_buffer — saved in state_dict, not optimized
train vs evalself.training flag, toggled by .train() / .eval()

58. Show it off

Concept

Out loud, slides closed: (1) trace the four forward-pass steps for a 3-sample batch; (2) state the dL/dx formula and why it has three terms; (3) explain why eval mode uses running stats instead of batch stats.

Stretch (homework): derive the full backward pass for batch normalization manually — every intermediate gradient. Implement LayerNorm from scratch and show it is equivalent to BatchNorm when batch_size=1. Train a 10-layer MLP with and without BatchNorm and plot the gradient norms per layer. Next up: dropout, regularization, and training stability — building on everything BatchNorm fixed.

59. Connect it up: Lesson 46: Batch Normalization

Connect it up

Draw it

One page, no notation unless you need it: draw how these connect — The forward pass · The backward pass · Running stats for inference · LayerNorm & GroupNorm · Your turn: build BatchNorm1d. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.

60. What you can do now

Recap

ideathe one thing to remember
forward pass4 steps: mu, var, x_hat = (x-mu)/std, y = gamma*x_hat+beta
backward dL/dx3-term chain: Ng - sum(g) - x_hatsum(gx_hat), all / (Nstd)
running statsEMA update each batch; used at eval (model.eval())
LayerNormsame formula, but over feature dim — no batch dependency

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 46 (Batch Normalization — forward, backward, running stats, LayerNorm, GroupNorm) — Barron · USAAIO Round 2 Preparation, 2026
  2. BatchNorm forward/backward, running stats, LayerNorm, GroupNorm all verified against nn.BatchNorm1d / nn.LayerNorm / nn.GroupNorm — torch 2.7.1 + numpy 2.2.6, real execution, June 2026

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108