summaryrefslogtreecommitdiff
path: root/ep_run/probe_cx3_rhogrid.py
blob: e154feca8a5683893a2de841e1dab0ebde4f044b (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
131
132
133
134
135
"""Dense rho(beta) trend probe (user: measure more divergence points, see the shape).
Trainer-faithful nudged relax (matches probe_rhorelax.py, validated to reproduce GOV meter),
K=30 sweeps to read the ASYMPTOTIC rho (not the 8-sweep transient), dense beta grid through and
past the ceiling, on MULTIPLE ckpts to see the ceiling sink with training. Reports per (ckpt,beta):
res0 (drive, should be proportional to beta if the loop is linear), asymptotic rho (tail-median of
res ratios), and the divergence verdict. GPU, read-only, no training."""
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='fw72m_plain2:35000,fw72m_plain2:95000,fw72m_plain2:150000,fw72m_plain2:230000')
ap.add_argument('--betas', default='0.03,0.06,0.1,0.15,0.2,0.3,0.4,0.5,0.6,0.7,0.8,1.0,1.3,1.7,2.5')
ap.add_argument('--K', type=int, default=30)
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 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))

betas = [float(s) for s in a.betas.split(',')]
first = True
for spec in a.ckpts.split(','):
    tag, step = spec.split(':'); step = int(step)
    p = f'runs/{tag}_s{step}.pt'
    try:
        ck = torch.load(p, map_location=dev, weights_only=False)
    except FileNotFoundError:
        print(f'{tag} 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)
        y = torch.stack([torch.from_numpy(data[i + 1:i + 1 + 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)
    W_out = ck['wout'].to(dev); ln_f = RMSNorm(C).to(dev); ln_f.load_state_dict(ck['lnf'])
    NBT = B * T
    def readout(z): return ln_f(z) @ W_out.t()
    print(f'\n=== {tag} s{step} (K={a.K} sweeps) ===', flush=True)
    print(f'{"beta":>7} {"res0":>10} {"res1":>10} {"resK":>11} {"rho_tail":>9} {"verdict":>10}', flush=True)
    bstar = None
    for beta in betas:
        f_ins, f_outs = [], []
        prev = tok(x).detach()
        for b in blocks:
            i = prev.detach().requires_grad_(True); o = b(i)
            f_ins.append(i); f_outs.append(o); prev = o.detach()
        ins, outs = f_ins, f_outs
        zs = [o.detach().float() for o in f_outs]; d = [None] * L
        res_list = []
        for k in range(a.K):
            zc = zs[L - 1].detach().requires_grad_(True)
            ce = F.cross_entropy(readout(zc).reshape(-1, vocab), y.reshape(-1))
            d[L - 1] = (-beta * NBT * torch.autograd.grad(ce, zc)[0]).detach().float()
            for l in range(L - 2, -1, -1):
                d[l] = torch.autograd.grad(outs[l + 1], ins[l + 1], grad_outputs=d[l + 1],
                                           retain_graph=True)[0].detach().float()
            prev = tok(x).detach(); n_ins, n_outs = [], []; rnum = rden = 0.0
            for l in range(L):
                i = prev.detach().requires_grad_(True); o = blocks[l](i)
                n_ins.append(i); n_outs.append(o)
                znew = o.detach().float() + d[l]
                rnum += float((znew - zs[l]).norm()); rden += float(zs[l].norm())
                zs[l] = znew; prev = zs[l]
            ins, outs = n_ins, n_outs
            res_list.append(rnum / max(rden, 1e-9))
            if not np.isfinite(res_list[-1]) or res_list[-1] > 1e4: break
        ratios = [res_list[i] / res_list[i - 1] for i in range(1, len(res_list))
                  if res_list[i - 1] > 1e-7]
        tail = ratios[-6:] if len(ratios) >= 6 else ratios
        rho_tail = float(np.median(tail)) if tail else float('nan')
        # resK-primary verdict: rho~1 AT THE FP FLOOR is noise/noise (documented artifact),
        # not divergence — rho only counts when the residual is materially off-floor
        diverged = (not np.isfinite(res_list[-1])) or res_list[-1] > 1e-2 or (rho_tail > 1.0 and res_list[-1] > 1e-4)
        verdict = 'DIVERGE' if diverged else 'converge'
        if diverged and bstar is None: bstar = beta
        print(f'{beta:>7.3f} {res_list[0]:>10.2e} {(res_list[1] if len(res_list)>1 else float("nan")):>10.2e} '
              f'{res_list[-1]:>11.2e} {rho_tail:>9.4f} {verdict:>10}', flush=True)
    print(f'  -> ceiling beta* (first DIVERGE) = {bstar}', flush=True)
    del tok, blocks, W_out, ln_f; torch.cuda.empty_cache()
print('\nRHOGRID_DONE', flush=True)