summaryrefslogtreecommitdiff
path: root/ep_run
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-07-09 04:31:59 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-07-09 04:31:59 -0500
commit57ad5f340a5148417d5f50420af0f7982ddc8105 (patch)
tree03e3bbfbf754de87d80ac05c0ccfbc87d490d3a0 /ep_run
parentcf96818498898575289a9430ad326d9450d3f3a0 (diff)
cascade: zil interleaved scheme = exact gradients at all depths (cos 1.0000 L6-24, trajectory 0.9998+); casc_ep_train zil trainer; naive-relaxation depth failure documented
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/casc_ep_train.py133
-rw-r--r--ep_run/cascade_probe.py57
2 files changed, 178 insertions, 12 deletions
diff --git a/ep_run/casc_ep_train.py b/ep_run/casc_ep_train.py
new file mode 100644
index 0000000..a5a5286
--- /dev/null
+++ b/ep_run/casc_ep_train.py
@@ -0,0 +1,133 @@
+"""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
diff --git a/ep_run/cascade_probe.py b/ep_run/cascade_probe.py
index 2f1449a..da3d380 100644
--- a/ep_run/cascade_probe.py
+++ b/ep_run/cascade_probe.py
@@ -18,11 +18,12 @@ ap.add_argument('--H', type=int, default=4); ap.add_argument('--T', type=int, de
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'], default='jacobi')
+ap.add_argument('--scheme', choices=['jacobi', 'gsf', 'gsr', 'zil'], 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)
args = ap.parse_args()
torch.manual_seed(args.seed)
dev = 'cuda' if torch.cuda.is_available() else 'cpu'
@@ -73,14 +74,15 @@ def fwd_states(z0):
return zs
def local_grad(zs, z0, l, beta, y):
- """dF/dz_l with neighbors fixed (zs: list of L free states, 0-indexed block l)."""
+ """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() / NBT
+ 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() / NBT
+ obj = obj + 0.5 * ((zs[l + 1].detach() - blocks[l + 1](zl, mask)) ** 2).sum() * sc
else:
- obj = obj + beta * F.cross_entropy(readout(zl).reshape(-1, vocab), y.reshape(-1))
+ obj = obj + beta * (NBT * sc) * F.cross_entropy(readout(zl).reshape(-1, vocab), y.reshape(-1))
return torch.autograd.grad(obj, zl)[0]
def relax(z0, y, beta):
@@ -106,14 +108,42 @@ def relax(z0, y, beta):
zs[l] = (zs[l] - args.eta * bufs[l]).detach()
return zs
-def dFdtheta(zs, z0, y, beta, params):
+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)."""
- E = 0.0; prev = z0
+ 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 []))
@@ -130,11 +160,14 @@ for bi in range(args.batches):
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]
- # two-phase
- zp = relax(z0, y, +args.beta); zm = relax(z0, y, -args.beta)
- gp = dFdtheta(zp, z0, y, +args.beta, gate_params)
- gm = dFdtheta(zm, z0, y, -args.beta, gate_params)
- gep = [(a - b) / (2 * args.beta) for a, b in zip(gp, gm)]
+ # two-phase (or single-phase interleaved zil)
+ 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)]
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):