diff options
| author | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 04:55:57 -0500 |
|---|---|---|
| committer | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 04:55:57 -0500 |
| commit | aa20066cbec3546236da55ee111d7265f8e6f0d7 (patch) | |
| tree | b89c37c887a707c15a2fc254c5f68cd15030ba54 /ep_run/casc_eq_train.py | |
| parent | 6ff624633ced5ff298612a8934de483a114e783e (diff) | |
casc_eq_train: equilibrium-mode cascade-EP trainer (GS-reverse solver K sweeps, two-phase ±β, EP readout at relaxed states) — the true-EP route C1
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run/casc_eq_train.py')
| -rw-r--r-- | ep_run/casc_eq_train.py | 140 |
1 files changed, 140 insertions, 0 deletions
diff --git a/ep_run/casc_eq_train.py b/ep_run/casc_eq_train.py new file mode 100644 index 0000000..d2da0ce --- /dev/null +++ b/ep_run/casc_eq_train.py @@ -0,0 +1,140 @@ +"""Cascade-EP trainer — EQUILIBRIUM MODE (the true-EP route). +Two-phase (+-beta) relaxation of all layer states to the nudged equilibria via +Gauss-Seidel reverse sweeps (solver choice only; readout is taken AT the relaxed +states with the standard EP formula), weight grad = (1/2beta)[dF/dtheta|+ - dF/dtheta|-]. +Inference = plain forward (standard LLM). Twin of casc_bp_train.py (same seed/data).""" +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('--tag', default='casc_eq6') +ap.add_argument('--L', type=int, default=6); ap.add_argument('--C', type=int, default=256) +ap.add_argument('--H', type=int, default=8); ap.add_argument('--T', type=int, default=256) +ap.add_argument('--B', type=int, default=24); ap.add_argument('--steps', type=int, default=4000) +ap.add_argument('--lr', type=float, default=3e-4); ap.add_argument('--warmup', type=int, default=200) +ap.add_argument('--beta', type=float, default=0.03); ap.add_argument('--seed', type=int, default=0) +ap.add_argument('--K', type=int, default=20) # GS reverse sweeps per phase +ap.add_argument('--geta', type=float, default=0.8) # GS state step (sum units) +ap.add_argument('--save_every', type=int, default=1000); ap.add_argument('--log', type=int, default=100) +ap.add_argument('--wandb', default=''); ap.add_argument('--wandb_run', default='') +args = ap.parse_args() +torch.manual_seed(args.seed) +dev = 'cuda' if torch.cuda.is_available() else 'cpu' + +DD = Path('/home/yurenh2/ept/ep_run/data/tinystories_bpe') +vocab = pickle.load(open(DD / 'meta.pkl', 'rb'))['vocab_size'] + +def get_batch(split): + data = np.memmap(DD / ('train.bin' if split == 'train' else '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]) + y = torch.stack([torch.from_numpy(data[i + 1:i + 1 + args.T].astype(np.int64)) for i in ix]) + return x.to(dev), y.to(dev) + +class Block(nn.Module): + 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) +mask = torch.triu(torch.full((args.T, args.T), float('-inf'), device=dev), 1) +readout = lambda z: z @ tok.weight.t() +all_params = list(tok.parameters()) + list(pos.parameters()) + list(blocks.parameters()) +opt = torch.optim.AdamW(all_params, lr=args.lr, weight_decay=1e-4) +sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / max(args.warmup, 1))) +NBT = args.B * args.T + +def free_states(x): + with torch.no_grad(): + z = tok(x) + pos(torch.arange(args.T, device=dev))[None] + z0 = z.clone(); zs = [] + for b in blocks: + z = b(z, mask); zs.append(z) + return z0, zs + +def relax(z0, zs_free, y, beta): + """GS reverse sweeps to the nudged equilibrium (SUM-unit local grads, step geta).""" + zs = [z.clone() for z in zs_free] + for _ in range(args.K): + for l in range(args.L - 1, -1, -1): + 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() + if l + 1 < args.L: + obj = obj + 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum() + else: + obj = obj + beta * NBT * F.cross_entropy(readout(zl).reshape(-1, vocab), y.reshape(-1)) + g = torch.autograd.grad(obj, zl)[0] + zs[l] = (zl - args.geta * g).detach() + return zs + +def dFdtheta(zs, x, y, beta): + """dF/dtheta at fixed relaxed states (z0 rebuilt WITH graph so emb gets its E-path grad).""" + prev = tok(x) + pos(torch.arange(args.T, device=dev))[None] + 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, all_params, allow_unused=True) + return [g if g is not None else None for g in gs] + +def ep_step(x, y): + z0, zs_free = free_states(x) + free_ce = F.cross_entropy(readout(zs_free[-1]).reshape(-1, vocab), y.reshape(-1)).item() + zp = relax(z0, zs_free, y, +args.beta) + zm = relax(z0, zs_free, y, -args.beta) + gp = dFdtheta(zp, x, y, +args.beta) + gm = dFdtheta(zm, x, y, -args.beta) + for p, a, b in zip(all_params, gp, gm): + p.grad = None if (a is None or b is None) else (a - b) / (2 * args.beta) + return free_ce + +@torch.no_grad() +def evaluate(nb=6): + tot = 0.0 + for _ in range(nb): + x, y = get_batch('val') + z = tok(x) + pos(torch.arange(args.T, device=dev))[None] + for b in blocks: z = b(z, mask) + tot += F.cross_entropy(readout(z).reshape(-1, vocab), y.reshape(-1)).item() + return tot / nb + +wb = None +if args.wandb: + try: + import wandb as _w + wb = _w.init(project=args.wandb, name=args.wandb_run or args.tag, id=args.wandb_run or args.tag, + resume='allow', config=vars(args)) + except Exception as e: + print(f'[wandb] disabled ({e})', flush=True) + +n = sum(p.numel() for p in all_params) +print(f'[{args.tag}] cascade-EP(EQUILIBRIUM) L{args.L} C{args.C} T{args.T} beta={args.beta} ' + f'K={args.K} geta={args.geta} | {n/1e6:.2f}M | {dev}', flush=True) +best, t0 = 1e9, time.time() +for step in range(args.steps + 1): + x, y = get_batch('train') + ce = ep_step(x, y) + torch.nn.utils.clip_grad_norm_(all_params, 1.0) + opt.step(); sched.step(); opt.zero_grad(set_to_none=True) + if step % args.log == 0: + val = evaluate(); best = min(best, val) + print(f'step {step:5d}/{args.steps} | train {ce:.4f} val {val:.4f} (best {best:.4f}) ' + f'| {step/max(time.time()-t0,1e-9):.3f} it/s', flush=True) + if wb is not None: + try: wb.log({'train_ce': ce, 'val_ce': val, 'best': best}, step=step) + except Exception: pass + if step % args.save_every == 0 and step > 0: + torch.save({'tok': tok.state_dict(), 'pos': pos.state_dict(), 'blocks': blocks.state_dict(), + 'step': step, 'val': best, 'config': vars(args)}, Path('runs') / f'{args.tag}_s{step}.pt') +print(f'[{args.tag}] DONE best val CE {best:.4f} (BP twin 2.9746; zil-diagnostic 3.3236)', flush=True) +if wb is not None: + try: wb.summary['best_val_ce'] = best; wb.finish() + except Exception: pass |
