summaryrefslogtreecommitdiff
path: root/ep_run/probe_bias.py
blob: 898f8f6356f07313dd62023f876010780cef864f (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
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
"""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

SEL = {'qkv_b0': names.index('0.attn.qkv.weight'), 'qkv_b8': names.index('8.attn.qkv.weight'),
       'w2_b6': names.index('6.ff.w2.weight'), 'w2_b11': names.index('11.ff.w2.weight')}
betas = [float(s) for s in a.betas.split(',')]

# self-check: persample with all +1 == plain
g0 = torch.Generator().manual_seed(99)
x, y = get_batch(g0)
_, gA = estimators(x, y, 3e-3, torch.ones(B, dtype=torch.long, device=dev))
_, gB = estimators(x, y, 3e-3, torch.ones(B, dtype=torch.long, device=dev))
same = max(float((a_ - b_).abs().max()) for a_, b_ in zip(gA, gB))
print(f'selfcheck determinism max|diff|={same:.2e} (must be ~0)', flush=True)

gbatch = torch.Generator().manual_seed(1234)
sgen = torch.Generator(device='cpu').manual_seed(777)
batches = [get_batch(gbatch) for _ in range(a.nb)]
for beta in betas:
    acc = {k: {'plain': [], 'psign': [], 'bp': []} for k in SEL}
    for (x, y) in batches:
        ones = torch.ones(B, dtype=torch.long, device=dev)
        s = (torch.randint(0, 2, (B,), generator=sgen) * 2 - 1).to(dev)
        g_bp, g_pl = estimators(x, y, beta, ones)
        _, g_ps = estimators(x, y, beta, s)
        for k, i in SEL.items():
            acc[k]['plain'].append((g_pl[i] - g_bp[i]).detach().cpu())
            acc[k]['psign'].append((g_ps[i] - g_bp[i]).detach().cpu())
            acc[k]['bp'].append(g_bp[i].detach().cpu())
        del g_bp, g_pl, g_ps
        torch.cuda.empty_cache()
    print(f'\n=== beta={beta:g}  (nb={a.nb}; rel_bias = |mean delta| / |mean g_bp|; sem = noise floor) ===', flush=True)
    print(f'{"layer":>8} {"plain rel_bias":>15} {"plain sem":>10} {"psign rel_bias":>15} {"psign sem":>10}', flush=True)
    for k in SEL:
        gb = torch.stack(acc[k]['bp']).mean(0); gnorm = float(gb.norm())
        out = []
        for est in ('plain', 'psign'):
            D = torch.stack(acc[k][est])
            mean = D.mean(0); rel = float(mean.norm()) / max(gnorm, 1e-30)
            sem = float((D - mean).pow(2).sum(dim=(1, 2)).mean().sqrt()) / (a.nb ** 0.5) / max(gnorm, 1e-30)
            out += [rel, sem]
        print(f'{k:>8} {out[0]:>15.4f} {out[1]:>10.4f} {out[2]:>15.4f} {out[3]:>10.4f}', flush=True)
print('\nBIAS_DONE', flush=True)