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
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).
Objectives
nn.BatchNorm1dBatchNorm1d from scratch including training and eval modesWarm-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.
Section
Part 1 of 4
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.
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.
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.
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.
Ranking
Put in order
Put the moves of Forward pass trace: 3 samples, 1 feature into the order they have to happen.
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.
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).
| sample | x | x - mu | x_hat | y = 2*x_hat - 1 |
|---|---|---|---|---|
| 1 | 1.0 | -3.0 | -1.2247 | -3.4495 |
| 2 | 4.0 | 0.0 | 0.0000 | -1.0000 |
| 3 | 7.0 | +3.0 | +1.2247 | +1.4495 |
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.
| sample | x | x - mu | x_hat | y = 2*x_hat - 1 |
|---|---|---|---|---|
| 1 | 1.0 | -3.0 | -1.2247 | -3.4495 |
| 2 | 4.0 | 0.0 | 0.0000 | -1.0000 |
| 3 | 7.0 | +3.0 | +1.2247 | +1.4495 |
Section
Part 2 of 4
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.
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.
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.
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.
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 = ...
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.
| gradient | autograd | manual formula | match |
|---|---|---|---|
| dL/dgamma[0] | -0.2690 | sum(dout*x_hat, axis=0)[0] | True |
| dL/dbeta[0] | -3.2438 | sum(dout, axis=0)[0] | True |
| dL/dx[0,0] | 0.2548 | 1/N * 1/std * (Ng - sum_g - x_hatsum_gx) | True |
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.
| gradient | autograd | manual formula | match |
|---|---|---|---|
| dL/dgamma[0] | -0.2690 | sum(dout*x_hat, axis=0)[0] | True |
| dL/dbeta[0] | -3.2438 | sum(dout, axis=0)[0] | True |
| dL/dx[0,0] | 0.2548 | 1/N * 1/std * (Ng - sum_g - x_hatsum_gx) | True |
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.
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.
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.
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.
Section
Part 3 of 4
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}} \]
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.
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.
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.
| batch | batch_mu | run_mean after |
|---|---|---|
| 1 | 4.0000 | 0.4000 |
| 2 | 5.0000 | 0.8600 |
| 3 | 6.0000 | 1.3740 |
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?
Concept
| mode | normalizes using | when used | call |
|---|---|---|---|
| train | current batch mu, sigma^2 | parameter updates | model.train() |
| eval | running mu, sigma^2 | inference / single sample | model.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.
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.
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.
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.
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.
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.
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.Section
Part 4 of 4
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.
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.
Concept
| norm | normalizes over | typical use | batch-size dep? |
|---|---|---|---|
| BatchNorm | batch dim (N) per feature | CNNs, MLPs | Yes — breaks at N=1 |
| LayerNorm | feature dim (D) per sample | Transformers, NLP | No |
| GroupNorm | groups of channels per sample | Small-batch vision | No |
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.
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.
| norm | normalizes over | typical use | batch-size dep? |
|---|---|---|---|
| BatchNorm | batch dim (N) per feature | CNNs, MLPs | Yes — breaks at N=1 |
| LayerNorm | feature dim (D) per sample | Transformers, NLP | No |
| GroupNorm | groups of channels per sample | Small-batch vision | No |
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.
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.
Ranking
Put in order
These are the steps of The BatchNorm recipe, scrambled. Put them back in order before the next slide shows you.
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.
Pattern
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:
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.
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.
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)?
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.
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.
Check
Which gradient is easiest and which is hardest?
Check your understanding
Which formula correctly gives dL/dbeta in BatchNorm?
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.
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.
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.
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?
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.
Section
Project
Concept
Implement BatchNorm1dScratch as an nn.Module with gamma, beta, running stats, and correct train/eval behavior. Verify every output against nn.BatchNorm1d.
| # | requirement | tool |
|---|---|---|
| 1 | forward pass (train mode): batch mu, var, x_hat, gamma*x_hat+beta | torch.mean, torch.var |
| 2 | running stats update (momentum=0.1) | register_buffer, manual EMA |
| 3 | forward pass (eval mode): use running stats | self.training flag |
| 4 | verify all outputs match nn.BatchNorm1d | torch.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.
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.
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| check | result |
|---|---|
| output shape matches input | True |
| x_hat mean per feature approx 0 | True |
| x_hat std per feature approx 1 | True |
| matches nn.BatchNorm1d output | True |
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))| comparison | result |
|---|---|
| output match (atol=1e-5) | True |
| running_mean match | True |
| scratch.training flag | True (default) |
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.
| comparison | result |
|---|---|
| output match (atol=1e-5) | True |
| running_mean match | True |
| scratch.training flag | True (default) |
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))| component | implementation detail |
|---|---|
| gamma / beta | nn.Parameter — learned by the optimizer |
| running_mean/var | register_buffer — saved in state_dict, not optimized |
| train vs eval | self.training flag, toggled by .train() / .eval() |
If match: True — you have implemented the same BatchNorm that every ResNet, BERT encoder, and GPT block uses.
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.
| component | implementation detail |
|---|---|
| gamma / beta | nn.Parameter — learned by the optimizer |
| running_mean/var | register_buffer — saved in state_dict, not optimized |
| train vs eval | self.training flag, toggled by .train() / .eval() |
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.
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.
Recap
nn.BatchNorm1d| idea | the one thing to remember |
|---|---|
| forward pass | 4 steps: mu, var, x_hat = (x-mu)/std, y = gamma*x_hat+beta |
| backward dL/dx | 3-term chain: Ng - sum(g) - x_hatsum(gx_hat), all / (Nstd) |
| running stats | EMA update each batch; used at eval (model.eval()) |
| LayerNorm | same formula, but over feature dim — no batch dependency |
Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.