summaryrefslogtreecommitdiff
path: root/ep_run/probe_statej.py
blob: f25325b73bd820d091bde2b9c012bc34a768e4fe (plain)
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
"""State-side Jacobian audit: per-block sigma(dO_l/dz_in) at the FREE operating state,
both lineages, matched steps. Weight-space audit (probe_specaudit) showed all weight scales
FLAT while the beta-normalized loop gain G=0.9/beta_cap grew ~300x -> hypothesis: the growth
is the CHAIN PRODUCT of per-block state Jacobians (each +tens-of-%, ^12 = hundreds-x).
Also dumps per-block max|logit| at the state (softcap calibration).
FD-JVP + autograd-vjp power iteration on J^T J (no forward-mode through SDPA)."""
import argparse, pickle
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from pathlib import Path

ap = argparse.ArgumentParser()
ap.add_argument('--ckpts', default=('plain:150000,plain:185000,plain:195000,plain:200000,'
                                    'plain:210000,plain:230000,cent:150000,cent:185000,'
                                    'cent:195000,cent:200000,cent:210000,cent:230000'))
ap.add_argument('--iters', type=int, default=25)
a = ap.parse_args()
dev = 'cuda'
torch.manual_seed(7)
B, T = 8, 256

class RMSNorm(nn.Module):
    def __init__(self, C, eps=1e-6):
        super().__init__(); self.g = nn.Parameter(torch.ones(C)); self.eps = eps
    def forward(self, x):
        return x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + self.eps) * self.g

class SwiGLU(nn.Module):
    def __init__(self, C):
        super().__init__()
        h = ((8 * C // 3) + 63) // 64 * 64
        self.w1 = nn.Linear(C, h, bias=False); self.w3 = nn.Linear(C, h, bias=False)
        self.w2 = nn.Linear(h, C, bias=False)
    def forward(self, x):
        return self.w2(F.silu(self.w1(x)) * self.w3(x))

class Olmo2Attn(nn.Module):
    def __init__(self, C, H, T):
        super().__init__()
        self.H, self.hd = H, C // H
        self.qkv = nn.Linear(C, 3 * C, bias=False); self.proj = nn.Linear(C, C, bias=False)
        self.qn, self.kn = RMSNorm(C), RMSNorm(C)
        inv = 1.0 / (500000.0 ** (torch.arange(0, self.hd, 2).float() / self.hd))
        fr = torch.outer(torch.arange(T).float(), inv)
        self.register_buffer('rc', fr.cos(), persistent=False)
        self.register_buffer('rs', fr.sin(), persistent=False)
    def rope(self, x):
        Tn = x.shape[2]
        x1, x2 = x[..., ::2], x[..., 1::2]
        c, s = self.rc[None, None, :Tn].to(x.dtype), self.rs[None, None, :Tn].to(x.dtype)
        return torch.stack((x1 * c - x2 * s, x1 * s + x2 * c), dim=-1).flatten(-2)
    def logits_stats(self, x):
        Bn, Tn, C = x.shape
        q, k, _ = self.qkv(x).split(C, dim=2)
        q, k = self.qn(q), self.kn(k)
        q = self.rope(q.view(Bn, Tn, self.H, self.hd).transpose(1, 2))
        k = self.rope(k.view(Bn, Tn, self.H, self.hd).transpose(1, 2))
        lg = (q @ k.transpose(-2, -1)) / (self.hd ** 0.5)
        mask = torch.ones(Tn, Tn, dtype=torch.bool, device=x.device).tril()
        lg = lg.masked_fill(~mask, 0.0)
        return float(lg.abs().max()), float(lg.abs().quantile(0.999))
    def forward(self, x):
        Bn, Tn, C = x.shape
        q, k, v = self.qkv(x).split(C, dim=2)
        q, k = self.qn(q), self.kn(k)
        q = self.rope(q.view(Bn, Tn, self.H, self.hd).transpose(1, 2))
        k = self.rope(k.view(Bn, Tn, self.H, self.hd).transpose(1, 2))
        v = v.view(Bn, Tn, self.H, self.hd).transpose(1, 2)
        y = F.scaled_dot_product_attention(q, k, v, is_causal=True)
        return self.proj(y.transpose(1, 2).contiguous().view(Bn, Tn, C))

class Olmo2Block(nn.Module):
    def __init__(self, C, H, T):
        super().__init__()
        self.attn = Olmo2Attn(C, H, T); self.ff = SwiGLU(C)
        self.na, self.nf = RMSNorm(C), RMSNorm(C)
    def forward(self, z):
        z = z + self.na(self.attn(z))
        return z + self.nf(self.ff(z))

def sigma_J(block, z0):
    """largest singular value of dblock(z)/dz at z0: power iteration on J^T J.
    Jv by central FD (fp32, scaled eps), J^T u by autograd."""
    z0 = z0.detach()
    v = torch.randn_like(z0); v /= v.norm()
    s = None
    for _ in range(a.iters):
        eps = 1e-3 * z0.norm() / (v.norm() * (z0.numel() ** 0.5) + 1e-30) * (z0.numel() ** 0.5)
        with torch.no_grad():
            jv = (block(z0 + eps * v) - block(z0 - eps * v)) / (2 * eps)
        zg = z0.clone().requires_grad_(True)
        o = block(zg)
        jtu = torch.autograd.grad((o * jv.detach()).sum(), zg)[0]
        s = float(jtu.norm().sqrt())          # |J^T J v|^{1/2} -> sigma as v converges
        v = jtu / (jtu.norm() + 1e-30)
    return s

first = True
for spec in a.ckpts.split(','):
    lineage, step = spec.split(':'); step = int(step)
    p = f'runs/fw72m_{lineage}_s{step}.pt'
    try:
        ck = torch.load(p, map_location='cpu', weights_only=False)
    except FileNotFoundError:
        print(f'{lineage} s{step}: MISSING', flush=True); continue
    cfg = ck['config']; C, H, L = cfg['C'], cfg['H'], cfg['L']
    if first:
        DD = Path('/home/yurenh2/ept/ep_run/data') / cfg.get('data', 'fineweb_edu')
        vocab = pickle.load(open(DD / 'meta.pkl', 'rb'))['vocab_size']
        data = np.memmap(DD / 'val.bin', dtype=np.uint16, mode='r')
        ix = torch.randint(len(data) - T - 1, (B,))
        x = torch.stack([torch.from_numpy(data[i:i + T].astype(np.int64)) for i in ix]).to(dev)
        first = False
    tok = nn.Embedding(vocab, C).to(dev); tok.load_state_dict(ck['tok'])
    blocks = nn.ModuleList([Olmo2Block(C, H, T) for _ in range(L)]).to(dev)
    blocks.load_state_dict(ck['blocks'], strict=False)
    with torch.no_grad():
        zs = [tok(x)]
        for b in blocks: zs.append(b(zs[-1]))
    sigs, lgs = [], []
    for l in range(L):
        sigs.append(sigma_J(blocks[l], zs[l]))
        lgs.append(blocks[l].attn.logits_stats(zs[l]))
    prod = float(np.prod(sigs))
    top = float(np.prod(sigs[6:]))
    print(f'{lineage} s{step//1000}k | sig_J per blk ' + ' '.join(f'{s:5.2f}' for s in sigs) +
          f' | PROD {prod:9.1f} top6 {top:7.1f} | max|logit| ' +
          ' '.join(f'{m:4.1f}' for m, _ in lgs), flush=True)
    del tok, blocks, zs
    torch.cuda.empty_cache()
print('STATEJ_DONE', flush=True)