"""Cascade-EP trainer (zil mode): train L DISTINCT standard transformer blocks with the interleaved reverse-sweep energy rule — per layer: update z_l (gamma=1, SUM units), then read that layer's theta-grad from its LOCAL energy term. Numerically == BP restructured as local two-factor rules; inference = plain forward (standard LLM). Twin of casc_bp_train.py (same seed => same init & data stream).""" 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_ep6') 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('--save_every', type=int, default=1000); ap.add_argument('--log', type=int, default=200) 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() io_params = list(tok.parameters()) + list(pos.parameters()) all_params = io_params + list(blocks.parameters()) blk_params = [list(b.parameters()) for b in blocks] 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 zil_step(x, y): """One cascade-EP (zil) gradient computation. Returns free-phase CE for logging; leaves grads in p.grad (normalized to CE_mean units).""" z0g = tok(x) + pos(torch.arange(args.T, device=dev))[None] z0 = z0g.detach() with torch.no_grad(): zs, z = [], z0 for b in blocks: z = b(z, mask); zs.append(z) free_ce = F.cross_entropy(readout(zs[-1]).reshape(-1, vocab), y.reshape(-1)).item() for p in all_params: p.grad = None for l in range(args.L - 1, -1, -1): zl = zs[l].detach().requires_grad_(True) # 1) state update: gamma=1 in SUM units (top gets the beta-nudge) if l + 1 < args.L: obj_u = 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum() else: obj_u = args.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() # 2) immediate local read: this block's params (+ io at the ends) prev = z0g if l == 0 else zs[l - 1].detach() # l=0 keeps emb graph params_l = blk_params[l] + (io_params if l == 0 else []) obj_r = 0.5 * ((zs[l] - blocks[l](prev, mask)) ** 2).sum() / NBT if l + 1 == args.L: obj_r = obj_r + args.beta * F.cross_entropy(readout(zs[l]).reshape(-1, vocab), y.reshape(-1)) params_l = params_l + [tok.weight] seen, uniq = set(), [] for p in params_l: if id(p) not in seen: seen.add(id(p)); uniq.append(p) gs = torch.autograd.grad(obj_r, uniq, allow_unused=True) for p, gg in zip(uniq, gs): if gg is None: continue p.grad = gg if p.grad is None else p.grad + gg for p in all_params: if p.grad is not None: p.grad /= 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(zil) L{args.L} C{args.C} H{args.H} T{args.T} beta={args.beta} | {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 = zil_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):.2f} 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 reference: casc_bp6)', flush=True) if wb is not None: try: wb.summary['best_val_ce'] = best; wb.finish() except Exception: pass