summaryrefslogtreecommitdiff
path: root/ep_run/cascade_probe.py
blob: a6bb3aac0dd22bb3de300202690132e300904ea4 (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
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
"""Cascade-EP gradient gate v2 — L distinct standard transformer blocks, layered energy
E = sum_l 0.5||z_l - f_l(z_{l-1})||^2.  Free equilibrium == the standard forward pass.
Two-phase (+-beta) relaxation of states -> local weight gradient; gate vs true BP.

v2: relaxation schemes (jacobi = physical parallel | gsf/gsr = Gauss-Seidel fwd/reverse sweeps),
    --include_io (emb/pos/tied-readout grads via the full EP formula), multi-batch gates,
    machine-parseable SUMMARY line.  Usage examples:
      python3 cascade_probe.py --L 12 --K 480
      python3 cascade_probe.py --L 6 --scheme gsr --K 12 --include_io
"""
import argparse, math, pickle, time
import numpy as np, torch, torch.nn as nn, torch.nn.functional as F
from pathlib import Path

ap = argparse.ArgumentParser()
ap.add_argument('--L', type=int, default=3); ap.add_argument('--C', type=int, default=128)
ap.add_argument('--H', type=int, default=4); ap.add_argument('--T', type=int, default=32)
ap.add_argument('--B', type=int, default=4); ap.add_argument('--beta', type=float, default=0.03)
ap.add_argument('--K', type=int, default=400); ap.add_argument('--eta', type=float, default=0.3)
ap.add_argument('--mom', type=float, default=0.9)
ap.add_argument('--scheme', choices=['jacobi', 'gsf', 'gsr', 'zil', 'sub', 'fb'], default='jacobi')
ap.add_argument('--batches', type=int, default=4)
ap.add_argument('--include_io', action='store_true')   # gate emb/pos/tied-readout too
ap.add_argument('--seed', type=int, default=0); ap.add_argument('--warp', type=float, default=1.0)
ap.add_argument('--ckpt', type=str, default='')         # gate at a casc_bp_train checkpoint
ap.add_argument('--sumscale', action='store_true')      # relax in SUM units (gamma=1 Z-IL semantics)
ap.add_argument('--init_sweep', action='store_true')    # one reverse gamma=1 sweep as state INIT (readout still at equilibrium = clean EP)
ap.add_argument('--sopt', choices=['sgd', 'adam'], default='sgd')   # state-relaxation optimizer
ap.add_argument('--prec', choices=['fp32', 'tf32', 'bf16'], default='fp32')  # A0.4 precision gate (EP side only; BP ref stays fp32)
ap.add_argument('--geta_auto', type=float, default=0.0)  # >0: per-layer gamma_l = c/(1+sigma_l+1^2), c=this; sigma via power-iter
args = ap.parse_args()
torch.manual_seed(args.seed)
dev = 'cuda' if torch.cuda.is_available() else 'cpu'
if args.prec == 'fp32':
    torch.set_float32_matmul_precision('highest')
elif args.prec == 'tf32':
    torch.backends.cuda.matmul.allow_tf32 = True
    torch.backends.cudnn.allow_tf32 = True
    torch.set_float32_matmul_precision('high')
_BF16 = (args.prec == 'bf16')

DD = Path('/home/yurenh2/ept/ep_run/data/tinystories_bpe')
vocab = pickle.load(open(DD / 'meta.pkl', 'rb'))['vocab_size']
data = np.memmap(DD / 'val.bin', dtype=np.uint16, mode='r')

class Block(nn.Module):
    """Standard pre-LN transformer block (distinct weights per layer, no recurrence)."""
    def __init__(self, C, H):
        super().__init__()
        self.ln1, self.ln2 = nn.LayerNorm(C), nn.LayerNorm(C)
        self.attn = nn.MultiheadAttention(C, H, batch_first=True)
        self.ff = nn.Sequential(nn.Linear(C, 4 * C), nn.GELU(), nn.Linear(4 * C, C))
    def forward(self, z, mask):
        h = self.ln1(z); a, _ = self.attn(h, h, h, attn_mask=mask, need_weights=False)
        z = z + a; return z + self.ff(self.ln2(z))

tok = nn.Embedding(vocab, args.C).to(dev)
pos = nn.Embedding(args.T, args.C).to(dev)
blocks = nn.ModuleList([Block(args.C, args.H) for _ in range(args.L)]).to(dev)
if args.ckpt:
    sd = torch.load(args.ckpt, map_location=dev)
    tok.load_state_dict(sd['tok']); pos.load_state_dict(sd['pos']); blocks.load_state_dict(sd['blocks'])
    print(f"[gate] loaded {args.ckpt} (step {sd.get('step')}, val {sd.get('val', float('nan')):.4f})")
elif args.warp != 1.0:
    with torch.no_grad():
        for p in blocks.parameters(): p.mul_(args.warp)
mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1)
readout = lambda z: z @ tok.weight.t()
io_params = list(tok.parameters()) + list(pos.parameters())
blk_params = list(blocks.parameters())
gate_params = blk_params + (io_params if args.include_io else [])
NBT = args.B * args.T

def get_batch(i):
    g = torch.Generator().manual_seed(1000 + i)
    ix = torch.randint(len(data) - args.T - 1, (args.B,), generator=g)
    x = torch.stack([torch.from_numpy(data[j:j + args.T].astype(np.int64)) for j in ix]).to(dev)
    y = torch.stack([torch.from_numpy(data[j + 1:j + 1 + args.T].astype(np.int64)) for j in ix]).to(dev)
    return x, y

def fwd_states(z0):
    zs = [z0]
    for b in blocks: zs.append(b(zs[-1], mask))
    return zs

def local_grad(zs, z0, l, beta, y):
    """dF/dz_l with neighbors fixed. sumscale: SUM-unit energy (gamma=1 == Z-IL exact sweep)."""
    sc = 1.0 if args.sumscale else 1.0 / NBT
    zl = zs[l].detach().requires_grad_(True)
    prev = z0 if l == 0 else zs[l - 1].detach()
    obj = 0.5 * ((zl - blocks[l](prev, mask)) ** 2).sum() * sc
    if l + 1 < args.L:
        obj = obj + 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum() * sc
    else:
        obj = obj + beta * (NBT * sc) * F.cross_entropy(readout(zl).reshape(-1, vocab), y.reshape(-1))
    return torch.autograd.grad(obj, zl)[0]

def layer_sigmas(z0, zs):
    """top sigma of J(blocks[l+1]) at zs[l] via power iteration; Jv by FD (SDPA-safe), J^T u by vjp."""
    sigs = []
    for l in range(args.L):
        if l + 1 >= args.L: sigs.append(0.0); continue
        zin = zs[l].detach()
        fn = lambda z: blocks[l + 1](z, mask)
        v = torch.randn_like(zin); v /= v.norm()
        sig = 0.0
        for _ in range(3):
            eps = 1e-3 * zin.norm() / max(v.norm(), 1e-12)
            with torch.no_grad():
                u = (fn(zin + eps * v) - fn(zin - eps * v)) / (2 * eps)   # J v (FD)
            sig = float(u.norm().item())                                   # ||Jv||, v normalized
            zi = zin.requires_grad_(True) if not zin.requires_grad else zin
            zi = zin.detach().requires_grad_(True)
            w = torch.autograd.grad(fn(zi), zi, grad_outputs=u.detach(), retain_graph=False)[0]
            v = (w / max(w.norm(), 1e-12)).detach()
        sigs.append(sig)
    return sigs

def relax(z0, y, beta):
    """minimize F_beta over states; free init = forward pass. Returns relaxed states."""
    with torch.no_grad():
        zs = [z.detach().clone() for z in fwd_states(z0)[1:]]
    gammas = None
    if args.geta_auto > 0:
        sigs = layer_sigmas(z0, zs)
        gammas = [args.geta_auto / (1.0 + s * s) for s in sigs]
    if args.init_sweep:                                  # numerical warm-start of the STATE only
        sv = args.sumscale; args.sumscale = True
        for l in range(args.L - 1, -1, -1):
            g = local_grad(zs, z0, l, beta, y)
            zs[l] = (zs[l] - g).detach()
        args.sumscale = sv
    if args.scheme == 'jacobi':
        for z in zs: z.requires_grad_(True)
        opt = (torch.optim.Adam(zs, lr=args.eta) if args.sopt == 'adam'
               else torch.optim.SGD(zs, lr=args.eta, momentum=args.mom))
        for _ in range(args.K):
            opt.zero_grad()
            E = 0.0; prev = z0
            for z, b in zip(zs, blocks): E = E + 0.5 * ((z - b(prev, mask)) ** 2).sum(); prev = z
            (E / NBT + beta * F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1))).backward()
            opt.step()
        return [z.detach() for z in zs]
    if args.scheme == 'fb':
        # forward-backward alternation (message passing): backward refreshes feedback
        # d_l = J_{l+1}^T d_{l+1} (top: -beta*NBT*dCE) at CURRENT states; forward REBUILDS
        # z_l = f_l(z_{l-1}) + d_l bottom-up. Converges O(beta)-fast; K rounds.
        d = [None] * args.L
        for _ in range(args.K):
            zc = zs[args.L - 1].detach().requires_grad_(True)
            ce = F.cross_entropy(readout(zc).reshape(-1, vocab), y.reshape(-1))
            d[args.L - 1] = (-beta * NBT * torch.autograd.grad(ce, zc)[0]).detach()
            for l in range(args.L - 2, -1, -1):
                zc = zs[l].detach().requires_grad_(True)
                fnext = blocks[l + 1](zc, mask)
                d[l] = torch.autograd.grad(fnext, zc, grad_outputs=d[l + 1])[0].detach()
            with torch.no_grad():
                alpha = args.eta if args.eta < 1.0 else 1.0   # damping mix (eta<1 => damped fb)
                prev = z0
                for l in range(args.L):
                    rebuilt = blocks[l](prev, mask) + d[l]
                    zs[l] = (1 - alpha) * zs[l] + alpha * rebuilt
                    prev = zs[l]
        return zs
    if args.scheme == 'sub':
        # assignment-form reverse sweeps: z_l := f_l(z_{l-1}) + J_{l+1}^T e_{l+1}
        # (top: z_L := f_L(z_{L-1}) - beta*NBT*dCE/dz_L). Unconditionally stable for small beta.
        for _ in range(args.K):
            for l in range(args.L - 1, -1, -1):
                prev = z0 if l == 0 else zs[l - 1].detach()
                with torch.no_grad():
                    ff = blocks[l](prev, mask)
                if l + 1 == args.L:
                    zc = zs[l].detach().requires_grad_(True)
                    ce = F.cross_entropy(readout(zc).reshape(-1, vocab), y.reshape(-1))
                    gce = torch.autograd.grad(ce, zc)[0]
                    zs[l] = (ff - beta * NBT * gce).detach()
                else:
                    zc = zs[l].detach().requires_grad_(True)
                    fnext = blocks[l + 1](zc, mask)
                    e_next = (zs[l + 1].detach() - fnext).detach()
                    jTe = torch.autograd.grad(fnext, zc, grad_outputs=e_next)[0]
                    zs[l] = (ff + jTe).detach()
        return zs
    order = range(args.L) if args.scheme == 'gsf' else range(args.L - 1, -1, -1)
    bufs = [torch.zeros_like(z) for z in zs]
    for _ in range(args.K):
        for l in order:
            g = local_grad(zs, z0, l, beta, y)
            bufs[l].mul_(args.mom).add_(g)
            step_l = (gammas[l] if gammas is not None else args.eta)
            zs[l] = (zs[l] - step_l * bufs[l]).detach()
    return zs

def dFdtheta(zs, z0, y, beta, params, x=None):
    """dF/dtheta at FIXED states (E-term local reads + beta*CE readout term for io params)."""
    prev = (tok(x) + pos(torch.arange(args.T, device=dev))[None]) if (args.include_io and x is not None) else z0
    E = 0.0
    for z, b in zip(zs, blocks): E = E + 0.5 * ((z - b(prev, mask)) ** 2).sum(); prev = z
    obj = E / NBT + beta * F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1))
    gs = torch.autograd.grad(obj, params, allow_unused=True)
    return [g if g is not None else torch.zeros(1, device=dev) for g in gs]

def zil_grads(x, y, beta):
    """Interleaved reverse sweep (Z-IL): update z_l (gamma=1, SUM units) then read theta_l
    IMMEDIATELY (e_l = -beta*delta_l exact at the feedforward point). Single phase; /beta.
    Returns grads matching gate_params, normalized like dFdtheta (E/NBT + CE_mean)."""
    z0g = tok(x) + pos(torch.arange(args.T, device=dev))[None]
    z0 = z0g.detach()
    with torch.no_grad():
        zs = [z.detach() for z in fwd_states(z0)[1:]]
    acc = {id(p): torch.zeros_like(p) for p in gate_params}
    for l in range(args.L - 1, -1, -1):
        zl = zs[l].detach().requires_grad_(True)
        prevd = z0 if l == 0 else zs[l - 1].detach()
        if l + 1 < args.L:
            obj_u = 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum()
        else:
            obj_u = beta * NBT * F.cross_entropy(readout(zl).reshape(-1, vocab), y.reshape(-1))
        g = torch.autograd.grad(obj_u, zl)[0]
        zs[l] = (zl - g).detach()                                   # gamma=1, SUM units
        prev_read = (z0g if (l == 0 and args.include_io) else prevd)
        obj_r = 0.5 * ((zs[l] - blocks[l](prev_read, mask)) ** 2).sum() / NBT
        if l + 1 == args.L:
            obj_r = obj_r + beta * F.cross_entropy(readout(zs[l]).reshape(-1, vocab), y.reshape(-1))
        gs = torch.autograd.grad(obj_r, gate_params, allow_unused=True)
        for p, gg in zip(gate_params, gs):
            if gg is not None: acc[id(p)] += gg
    return [acc[id(p)] / beta for p in gate_params]

flat = lambda gs: torch.cat([g.reshape(-1) for g in gs])
names = ([f'blk.{n}' for n, _ in blocks.named_parameters()] +
         ([f'io.{n}' for n, _ in list(tok.named_parameters()) + list(pos.named_parameters())] if args.include_io else []))

cos_b, shr_b, blkcos = [], [], [[] for _ in range(args.L)]
t0 = time.time()
for bi in range(args.batches):
    x, y = get_batch(bi)
    z0 = (tok(x) + pos(torch.arange(args.T, device=dev))[None]).detach()
    # BP reference (embedding z0 grad flows for io gate)
    z0ref = tok(x) + pos(torch.arange(args.T, device=dev))[None]
    zs = [z0ref]
    for b in blocks: zs.append(b(zs[-1], mask))
    ce = F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1))
    gbp = list(torch.autograd.grad(ce, gate_params, allow_unused=True))
    gbp = [g if g is not None else torch.zeros(1, device=dev) for g in gbp]
    # EP side: pure-bf16 MODEL when --prec bf16 (BP reference above stays fp32); cast back after
    if _BF16:
        sd32 = [ (m, {k: v.clone() for k, v in m.state_dict().items()}) for m in [tok, pos, blocks] ]
        for m in [tok, pos, blocks]: m.bfloat16()
        z0 = z0.bfloat16(); mask = mask.bfloat16()
    if args.scheme == 'zil':
        gep = zil_grads(x, y, args.beta)
    else:
        zp = relax(z0, y, +args.beta); zm = relax(z0, y, -args.beta)
        gp = dFdtheta(zp, z0, y, +args.beta, gate_params, x)
        gm = dFdtheta(zm, z0, y, -args.beta, gate_params, x)
        gep = [(a - b) / (2 * args.beta) for a, b in zip(gp, gm)]
    gep = [g.float() for g in gep]
    if _BF16:
        for m, sd in sd32:
            m.float(); m.load_state_dict(sd)
        mask = mask.float()
    cos_b.append(F.cosine_similarity(flat(gep), flat(gbp), dim=0).item())
    shr_b.append((flat(gep).norm() / flat(gbp).norm()).item())
    for l in range(args.L):
        idx = [i for i, n in enumerate(names) if n.startswith(f'blk.{l}.')]
        blkcos[l].append(F.cosine_similarity(flat([gep[i] for i in idx]), flat([gbp[i] for i in idx]), dim=0).item())

cm, cmin = float(np.mean(cos_b)), float(np.min(cos_b))
sm = float(np.mean(shr_b))
print('per-block cos: ' + '  '.join(f'{l}:{np.mean(v):.4f}' for l, v in enumerate(blkcos)))
if args.include_io:
    idx = [i for i, n in enumerate(names) if n.startswith('io.')]
    print(f'io cos (last batch): {F.cosine_similarity(flat([gep[i] for i in idx]), flat([gbp[i] for i in idx]), dim=0).item():.4f}')
print(f'SUMMARY prec={args.prec} scheme={args.scheme}{"+init" if args.init_sweep else ""}+{args.sopt} L={args.L} C={args.C} K={args.K} eta={args.eta} beta={args.beta} '
      f'io={int(args.include_io)} warp={args.warp} ckpt={args.ckpt or "-"} '
      f'cos={cm:.4f} cosmin={cmin:.4f} shrink={sm:.3f} t={time.time()-t0:.1f}s')