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
|
"""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')
diverged = (not np.isfinite(res_list[-1])) or res_list[-1] > 1e-2 or rho_tail > 1.0
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)
|