diff options
| author | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 04:04:21 -0500 |
|---|---|---|
| committer | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 04:04:21 -0500 |
| commit | 603a32a4e8f3dea3b582a97d27b9c82e8daebfa4 (patch) | |
| tree | ab4d94d32b224f904e0834059995c074e7645b41 /ep_run | |
| parent | a015da4a613f71f23dcfe2168c7d783faa0a6c55 (diff) | |
cascade_probe: layered-energy EP on L distinct standard blocks — free equilibrium == plain forward (standard LLM inference), two-phase grad gate cos 0.9968 vs BP (L3 C128)
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run')
| -rw-r--r-- | ep_run/cascade_probe.py | 102 |
1 files changed, 102 insertions, 0 deletions
diff --git a/ep_run/cascade_probe.py b/ep_run/cascade_probe.py new file mode 100644 index 0000000..572097c --- /dev/null +++ b/ep_run/cascade_probe.py @@ -0,0 +1,102 @@ +"""Cascade-EP gradient gate: 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 +(inference = plain LLM). Training = two-phase (±beta nudge) relaxation of z, weight +grad = (1/2beta)[dE/dtheta|+ - dE/dtheta|-]. Gate: cosine vs true BP gradient.""" +import argparse, math, 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('--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('--seed', type=int, default=0); ap.add_argument('--trained_warp', type=float, default=1.0) +args = ap.parse_args() +torch.manual_seed(args.seed) +dev = 'cuda' if torch.cuda.is_available() else 'cpu' +torch.set_float32_matmul_precision('highest') + +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') +ix = torch.randint(len(data) - args.T - 1, (args.B,)) +x = torch.stack([torch.from_numpy(data[i:i + args.T].astype(np.int64)) for i in ix]).to(dev) +y = torch.stack([torch.from_numpy(data[i + 1:i + 1 + args.T].astype(np.int64)) for i in ix]).to(dev) + +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.trained_warp != 1.0: # crude "not-at-init" landscape + with torch.no_grad(): + for p in blocks.parameters(): p.mul_(args.trained_warp) +mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1) +readout = lambda z: z @ tok.weight.t() + +def fwd_states(z0): + zs = [z0] + for b in blocks: zs.append(b(zs[-1], mask)) + return zs + +z0 = (tok(x) + pos(torch.arange(args.T, device=dev))[None]).detach() + +# ---------------- true BP reference ---------------- +for p in blocks.parameters(): p.requires_grad_(True) +zs = fwd_states(z0) +ce = F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1)) +gbp = torch.autograd.grad(ce, list(blocks.parameters()), allow_unused=True) +gbp = [g if g is not None else torch.zeros(1, device=dev) for g in gbp] +print(f'BP ref: CE {ce.item():.4f}') + +# ---------------- two-phase cascade EP ---------------- +def relax(beta): + """minimize E + beta*CE over free states z_1..z_L by GD (free init = forward pass).""" + with torch.no_grad(): + zs = [z.detach().clone() for z in fwd_states(z0)[1:]] + for z in zs: z.requires_grad_(True) + opt = torch.optim.SGD(zs, lr=args.eta, momentum=0.9) + for k 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 + Fobj = E / (args.B * args.T) + beta * F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1)) + Fobj.backward() + opt.step() + return [z.detach() for z in zs], (E / (args.B * args.T)).item() + +def dEdtheta(zs): + """dE/dtheta at fixed states (the local EP readout).""" + for p in blocks.parameters(): + if p.grad is not None: p.grad = None + 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 / (args.B * args.T)).backward() + return [(p.grad.clone() if p.grad is not None else torch.zeros(1, device=dev)) for p in blocks.parameters()] + +zp, Ep = relax(+args.beta) +zm, Em = relax(-args.beta) +gp, gm = dEdtheta(zp), dEdtheta(zm) +gep = [(a - b) / (2 * args.beta) for a, b in zip(gp, gm)] +print(f'nudged relax E+: {Ep:.3e} E-: {Em:.3e} (K={args.K}, eta={args.eta}, beta={args.beta})') + +# ---------------- gate ---------------- +names = [n for n, _ in blocks.named_parameters()] +flat = lambda gs: torch.cat([g.reshape(-1) for g in gs]) +cos_all = F.cosine_similarity(flat(gep), flat(gbp), dim=0).item() +print(f'\nGATE cos(cascadeEP, BP) overall: {cos_all:.4f}') +for l in range(args.L): + idx = [i for i, n in enumerate(names) if n.startswith(f'{l}.')] + c = F.cosine_similarity(flat([gep[i] for i in idx]), flat([gbp[i] for i in idx]), dim=0).item() + r = (flat([gep[i] for i in idx]).norm() / flat([gbp[i] for i in idx]).norm()).item() + print(f' block {l}: cos {c:.4f} |EP|/|BP| {r:.3f}') |
