配套教程:归一化 / 残差 / 初始化 · EN
Normalization / Residual / Initialization - minimal runnable implementation
1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182
"""
Normalization / Residual / Initialization - minimal runnable implementation
===========================================================================
Educational PyTorch reference for normalization layers, the Pre-LN vs Post-LN
gradient behaviour, and variance-preserving initialization.
Standalone script: runs six sanity checks on CPU in a few seconds.
Pairs with: docs/tutorials/normalization_init_tutorial.md (concept reference).
What it demonstrates:
layernorm_from_scratch - LayerNorm matches torch.nn.LayerNorm
rmsnorm_from_scratch - RMSNorm matches torch.nn.RMSNorm, and (unlike LayerNorm)
is NOT mean-shift invariant (it only re-scales, no re-centering)
Sanity checks:
[a] LayerNorm from scratch == nn.LayerNorm (population var, eps inside the sqrt)
[b] RMSNorm from scratch == nn.RMSNorm; LayerNorm(x+c)==LayerNorm(x) but RMSNorm(x+c)!=RMSNorm(x)
[c] BatchNorm train != eval: train uses batch stats (+updates running stats), eval uses running stats
[d] Post-LN piles parameter gradients near the OUTPUT (top-heavy, last/first>1, needs warmup); Pre-LN piles them near the INPUT instead (bottom-heavy, opposite skew, comparable magnitude)
[e] Kaiming preserves the second moment E[y^2] through Linear+ReLU (factor 2/fan_in); Xavier-for-ReLU halves it to ~0.5
[f] GPT-2 residual scaling 1/sqrt(2N) keeps the residual-stream variance bounded vs linear growth
Run:
python normalization.py
"""
import torch
import torch.nn as nn
import torch.nn.functional as F
torch.manual_seed(0)
def layernorm_from_scratch(x, weight, bias, eps=1e-5):
"""LayerNorm over the last dim. x: [..., d]. Population variance (unbiased=False),
eps added INSIDE the sqrt, then affine: y = (x-mean)/sqrt(var+eps) * weight + bias."""
mean = x.mean(dim=-1, keepdim=True)
var = x.var(dim=-1, unbiased=False, keepdim=True) # population var, matches torch
return (x - mean) / torch.sqrt(var + eps) * weight + bias
def rmsnorm_from_scratch(x, weight, eps=1e-6):
"""RMSNorm over the last dim: only re-scale by the RMS, NO mean-centering, NO bias.
y = x / sqrt(mean(x^2) + eps) * weight."""
ms = x.pow(2).mean(dim=-1, keepdim=True) # mean of squares
return x / torch.sqrt(ms + eps) * weight
class ResidualStack(nn.Module):
"""A deep stack of identical residual blocks, either Pre-LN or Post-LN, plus a final
linear head, used to show how parameter-gradient norms distribute across depth at init.
Pre-LN block: h = h + Linear(LayerNorm(h))
Post-LN block: h = LayerNorm(h + Linear(h))
The head matters: without it, a Post-LN stack's output is LayerNorm-normalized, so a
loss like mean(out^2) is ~scale-invariant and artificially suppresses ALL gradients
(a confound). The head makes the loss depend on the actual output, so the only signal
left is the genuine cross-depth gradient distribution.
"""
def __init__(self, depth, d, pre_ln=True):
super().__init__()
self.pre_ln = pre_ln
self.lns = nn.ModuleList(nn.LayerNorm(d) for _ in range(depth))
self.lins = nn.ModuleList(nn.Linear(d, d) for _ in range(depth))
self.head = nn.Linear(d, d) # final projection (un-normalizes the output)
def forward(self, h):
for ln, lin in zip(self.lns, self.lins):
if self.pre_ln:
h = h + lin(ln(h))
else:
h = ln(h + lin(h))
return self.head(h)
def block_grad_topheavy(depth, d, pre_ln):
"""Build a fresh stack, one forward+backward, return the LAST-block / FIRST-block
weight-grad-norm ratio. >1 means the gradient is concentrated near the OUTPUT
(top-heavy). Using a relative ratio (rather than an absolute grad norm) is meant to
reduce sensitivity to the loss choice / output normalization; this script only runs
a single seed with a single quadratic loss, so "robust to loss choice" itself is not
exhaustively verified here (would need multi-seed stats and a non-quadratic loss)."""
torch.manual_seed(0) # same init for both variants
stack = ResidualStack(depth, d, pre_ln=pre_ln)
stack(torch.randn(16, d)).pow(2).mean().backward()
gn = [lin.weight.grad.norm().item() for lin in stack.lins]
return gn[-1] / gn[0]
def main():
B, d, eps = 8, 64, 1e-5
# [a] LayerNorm from scratch == nn.LayerNorm
x = torch.randn(B, d)
ln = nn.LayerNorm(d, eps=eps) # affine: weight=1, bias=0 at init
mine = layernorm_from_scratch(x, ln.weight, ln.bias, eps=eps)
a_ok = torch.allclose(mine, ln(x), atol=1e-5)
print(f"[a] LayerNorm from scratch vs nn.LayerNorm: max|Δ| = {(mine - ln(x)).abs().max():.2e} "
f"{'OK' if a_ok else 'FAIL'}")
assert a_ok
# [b] RMSNorm matches nn.RMSNorm; and RMSNorm is NOT mean-shift invariant (LayerNorm is)
rms = nn.RMSNorm(d, eps=1e-6)
mine_r = rmsnorm_from_scratch(x, rms.weight, eps=1e-6)
b1 = torch.allclose(mine_r, rms(x), atol=1e-5)
c = 5.0 # constant shift on every feature
ln_shift = (ln(x + c) - ln(x)).abs().max().item() # LayerNorm removes the mean -> ~0
rms_shift = (rmsnorm_from_scratch(x + c, rms.weight) - mine_r).abs().max().item() # RMS keeps it -> >0
b_ok = b1 and ln_shift < 1e-4 and rms_shift > 1e-2
print(f"[b] RMSNorm vs nn.RMSNorm: max|Δ| = {(mine_r - rms(x)).abs().max():.2e}; "
f"mean-shift LN |Δ|={ln_shift:.2e} (~0, re-centers) vs RMS |Δ|={rms_shift:.2e} (>0, only re-scales) "
f"{'OK' if b_ok else 'FAIL'}")
assert b_ok
# [c] BatchNorm: train uses batch stats (+ updates running stats), eval uses running stats
bn = nn.BatchNorm1d(d) # running_mean=0, running_var=1 at init
xb = torch.randn(B, d) * 3 + 7 # shifted/scaled so batch stats != running
bn.train()
y_train = bn(xb)
running_moved = bn.running_mean.abs().mean().item() # moved away from 0 toward batch mean
bn.eval()
y_eval = bn(xb)
differ = (y_train - y_eval).abs().max().item()
# train output is normalized by the batch -> ~0 mean / ~1 std per feature; eval is not
train_mean = y_train.mean(dim=0).abs().mean().item()
c_ok = running_moved > 0 and differ > 1e-2 and train_mean < 1e-4
print(f"[c] BatchNorm train!=eval: max|Δ| = {differ:.2e}; running_mean moved {running_moved:.3f} from 0; "
f"train per-feature mean = {train_mean:.2e} (~0) {'OK' if c_ok else 'FAIL'}")
assert c_ok
# [d] Post-LN concentrates parameter gradients near the OUTPUT (top-heavy -> needs warmup);
# Pre-LN's clean identity path instead skews slightly toward the INPUT (bottom-heavy),
# at a comparable magnitude but opposite direction (Xiong et al. 2020's asymptotic claim)
depth = 48
r_pre = block_grad_topheavy(depth, d, pre_ln=True) # ~0.4: bottom-heavy (bottom grad ~2.5x the top)
r_post = block_grad_topheavy(depth, d, pre_ln=False) # ~2.3: top-heavy (gradient piled at the top)
d_ok = r_post > 1.3 and r_post > 2 * r_pre # Post-LN top-heavy AND more imbalanced than Pre-LN
print(f"[d] per-block weight-grad top/bottom ratio (last/first over {depth} blocks): "
f"Pre-LN={r_pre:.2f} (bottom-heavy) Post-LN={r_post:.2f} (top-heavy, >1) "
f"-> opposite skew, comparable magnitude; Post-LN piles gradient near the output, needs warmup "
f"{'OK' if d_ok else 'FAIL'}")
assert d_ok
# [e] Kaiming preserves the SECOND MOMENT E[y^2] through Linear+ReLU (the quantity He et al.
# actually propagate): ReLU halves E[pre^2], so the factor 2/fan_in restores E[y^2]~E[x^2]=1.
# Xavier (1/fan_in) lacks the factor 2 -> E[y^2]~0.5, decaying by half every layer.
fan_in, fan_out, N = 512, 512, 4096
xin = torch.randn(N, fan_in) # E[x^2] = 1 (unit second moment)
W_kaiming = torch.randn(fan_out, fan_in) * (2.0 / fan_in) ** 0.5 # He: std = sqrt(2/fan_in)
W_xavier = torch.randn(fan_out, fan_in) * (1.0 / fan_in) ** 0.5 # Xavier: std = sqrt(1/fan_in)
ms_kaiming = F.relu(xin @ W_kaiming.t()).pow(2).mean().item() # E[y^2] ~ 1 (preserved)
ms_xavier = F.relu(xin @ W_xavier.t()).pow(2).mean().item() # E[y^2] ~ 0.5 (halved -> decays)
e_ok = abs(ms_kaiming - 1.0) < 0.1 and abs(ms_xavier - 0.5) < 0.1
print(f"[e] post-ReLU second moment E[y^2] (input E[x^2]=1): Kaiming = {ms_kaiming:.3f} (~1, preserved) "
f"Xavier = {ms_xavier:.3f} (~0.5, halves per layer) {'OK' if e_ok else 'FAIL'}")
assert e_ok
# [f] GPT-2 residual scaling 1/sqrt(2N): residual-stream variance stays bounded vs linear growth.
# Nlayers = number of transformer layers; each layer writes to the residual stream TWICE
# (attn output-proj + FFN down-proj), so there are 2*Nlayers total writes, each scaled by
# 1/sqrt(2*Nlayers) -- matching GPT-2's actual "N layers -> 2N writes" semantics.
Nlayers, width = 50, 256
n_writes = 2 * Nlayers
h_plain = torch.randn(N, width)
h_scaled = h_plain.clone()
v0 = h_plain.var().item()
for _ in range(n_writes):
delta = torch.randn(N, width) # each write adds an O(1)-variance update
h_plain = h_plain + delta # no scaling -> Var grows ~linearly with #writes
h_scaled = h_scaled + delta * (1.0 / n_writes) ** 0.5 # GPT-2 1/sqrt(2N) residual scaling
growth_plain = h_plain.var().item() / v0 # ~ (1 + 2*Nlayers)
growth_scaled = h_scaled.var().item() / v0 # ~ O(1), bounded
f_ok = growth_plain > 10 * growth_scaled
print(f"[f] residual-stream var growth over {Nlayers} layers ({n_writes} writes): "
f"unscaled ×{growth_plain:.1f} vs 1/sqrt(2N)-scaled ×{growth_scaled:.2f} "
f"{'OK' if f_ok else 'FAIL'}")
assert f_ok
print("\nall normalization / residual / init sanity checks passed ✓")
if __name__ == "__main__":
main()