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
|
"""Suspect-2 probe (C768 plateau): does Muon msign AMPLIFY the tiny EP-vs-BP discrepancy?
Per weight matrix: cos(gEP,gBP) raw vs cos(msign(gEP),msign(gBP)). Derived from probe_bias.
Old docstring: Direct measurement of the systematic single-sided EP bias (the audit gap: BBP probes
measured NOISE thoroughly, discarded the mean/bias component unreported) + verification of the
user's claim: PER-SAMPLE random nudge sign makes the batch-averaged estimator unbiased at O(beta).
Three estimators at fixed ckpt, N batches, two betas:
plain : all samples nudged +beta -> bias = beta*E[B] (systematic, batch-indep)
persample : sample i nudged s_i*beta, read re-flipped by s_i -> odd-order bias cancels
WITHIN each batch (sqrt(B) suppression per step, 0 in expectation)
BP : exact reference.
Self-check: persample with all s=+1 must equal plain exactly.
Report per layer family: |mean(ghat-gbp)| / |mean(gbp)| (systematic), across-batch sem (noise
floor), for beta in {3e-3, 1e-2} — plain's bias should scale ~x3.3, persample's should sit at
the noise floor. Read-only, coexists on GPU0."""
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('--ckpt', default='runs/fw72m_plain2_s35000.pt')
ap.add_argument('--K', type=int, default=3)
ap.add_argument('--betas', default='3e-3,1e-2')
ap.add_argument('--nb', type=int, default=16)
a = ap.parse_args()
dev = 'cuda'
torch.manual_seed(11)
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))
ck = torch.load(a.ckpt, map_location=dev, weights_only=False)
cfg = ck['config']; C, H, L = cfg['C'], cfg['H'], cfg['L']
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')
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
params = list(blocks.parameters())
names = [n for n, _ in blocks.named_parameters()]
def readout(z): return ln_f(z) @ W_out.t()
def get_batch(g):
ix = torch.randint(len(data) - T - 1, (B,), generator=g)
x = torch.stack([torch.from_numpy(data[i:i + T].astype(np.int64)) for i in ix])
y = torch.stack([torch.from_numpy(data[i + 1:i + 1 + T].astype(np.int64)) for i in ix])
return x.to(dev), y.to(dev)
def grads_from(outs_list, cots):
obj = sum((o * c.detach()).sum() for o, c in zip(outs_list, cots))
gs = torch.autograd.grad(obj, params, allow_unused=True, retain_graph=True)
return [g.float() if g is not None else torch.zeros_like(p) for g, p in zip(gs, params)]
def estimators(x, y, beta, signs):
"""signs: (B,) tensor of +-1. Returns (g_bp, g_ep) where g_ep uses per-sample sign nudges."""
z = tok(x)
zs_bp = []
for b in blocks:
z = b(z); zs_bp.append(z)
ce = F.cross_entropy(readout(zs_bp[-1]).reshape(-1, vocab), y.reshape(-1))
g_bp = [g.float() for g in torch.autograd.grad(ce, params, retain_graph=True)]
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
sB = signs.view(B, 1, 1).float()
for k in range(a.K):
zc = zs[L - 1].detach().requires_grad_(True)
# per-sample signed top objective: sum_i s_i * (sum_t loss_it) (token-sum, matches beta*NBT*mean scaling)
ce_tok = F.cross_entropy(readout(zc).reshape(-1, vocab), y.reshape(-1), reduction='none').view(B, T)
obj_top = (signs.float()[:, None] * ce_tok).sum()
d[L - 1] = (-beta * torch.autograd.grad(obj_top, 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 = [], []
for l in range(L):
i = prev.detach().requires_grad_(True)
o = blocks[l](i)
n_ins.append(i); n_outs.append(o)
zs[l] = o.detach().float() + d[l]
prev = zs[l]
ins, outs = n_ins, n_outs
md = [-(sB * di) / (beta * NBT) for di in d] # re-flip per sample; scale matches trainer
g_ep = grads_from(outs, md)
return g_bp, g_ep
from muon import newton_schulz
FAMS = {'qkv_bot': [f'{l}.attn.qkv.weight' for l in range(0,4)],
'qkv_top': [f'{l}.attn.qkv.weight' for l in range(8,12)],
'w2_bot': [f'{l}.ff.w2.weight' for l in range(0,4)],
'w2_top': [f'{l}.ff.w2.weight' for l in range(8,12)],
'proj_top':[f'{l}.attn.proj.weight' for l in range(8,12)]}
IDX = {k: [names.index(n) for n in v] for k, v in FAMS.items()}
g0 = torch.Generator().manual_seed(1234)
def cosf(a, b):
return float((a*b).sum() / (a.norm()*b.norm() + 1e-30))
acc = {k: {'raw': [], 'ms': []} for k in IDX}
for bi in range(6):
x, y = get_batch(g0)
ones = torch.ones(B, dtype=torch.long, device=dev)
g_bp, g_ep = estimators(x, y, 3e-3, ones)
for k, idxs in IDX.items():
for i in idxs:
acc[k]['raw'].append(cosf(g_ep[i], g_bp[i]))
me = newton_schulz(g_ep[i].bfloat16()).float()
mb = newton_schulz(g_bp[i].bfloat16()).float()
acc[k]['ms'].append(cosf(me, mb))
del g_bp, g_ep
torch.cuda.empty_cache()
import numpy as np
print(f'{"family":>9} {"cos_raw":>9} {"cos_msign":>10} (mean over 6 batches x 4 mats)')
for k in IDX:
print(f'{k:>9} {np.mean(acc[k]["raw"]):>9.4f} {np.mean(acc[k]["ms"]):>10.4f}')
print('MSIGN_DONE', flush=True)
|