⌂ 全部教程 easyaioffer.github.io ⭐ Star

normalization.py

配套教程:归一化 / 残差 / 初始化 · EN

Normalization / Residual / Initialization - minimal runnable implementation

在 GitHub 查看 原始 .py python normalization.py
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()