diff options
| author | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 07:44:14 -0500 |
|---|---|---|
| committer | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 07:44:14 -0500 |
| commit | ddff6bf7d31b0dd28bde1d057939c2c4f9b359ba (patch) | |
| tree | 64a9fa748b7fec6225d828694220b9f3cf3cb05a /ep_run | |
| parent | aa20066cbec3546236da55ee111d7265f8e6f0d7 (diff) | |
cascade equilibrium solver solved: fb (forward-backward message passing) K=3 beta=0.003 — gates 1.0000/0.9990/0.9946 across BP trajectory, L12 0.9999; trainer v2 single-sided EP readout + divergence guard (~5x BP)
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_eq_train.py | 55 | ||||
| -rw-r--r-- | ep_run/cascade_probe.py | 72 |
2 files changed, 105 insertions, 22 deletions
diff --git a/ep_run/casc_eq_train.py b/ep_run/casc_eq_train.py index d2da0ce..99ff35c 100644 --- a/ep_run/casc_eq_train.py +++ b/ep_run/casc_eq_train.py @@ -13,9 +13,9 @@ ap.add_argument('--L', type=int, default=6); ap.add_argument('--C', type=int, de 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('--beta', type=float, default=0.003); ap.add_argument('--seed', type=int, default=0) +ap.add_argument('--K', type=int, default=3) # fb (message-passing) rounds +ap.add_argument('--geta', type=float, default=1.0) # fb mixing (1.0 = undamped) 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() @@ -61,19 +61,24 @@ def free_states(x): return z0, zs def relax(z0, zs_free, y, beta): - """GS reverse sweeps to the nudged equilibrium (SUM-unit local grads, step geta).""" + """fb rounds to the nudged equilibrium: backward refresh of feedback d_l = J_{l+1}^T d_{l+1} + (top: -beta*NBT*dCE at current top), then forward REBUILD z_l = f_l(z_{l-1}) + d_l.""" zs = [z.clone() for z in zs_free] + d = [None] * args.L 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() + 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(): + prev = z0 + for l in range(args.L): + rebuilt = blocks[l](prev, mask) + d[l] + zs[l] = (1 - args.geta) * zs[l] + args.geta * rebuilt if args.geta < 1.0 else rebuilt + prev = zs[l] return zs def dFdtheta(zs, x, y, beta): @@ -86,14 +91,24 @@ def dFdtheta(zs, x, y, beta): return [g if g is not None else None for g in gs] def ep_step(x, y): + """single-sided EP: relax to the +beta equilibrium; grad = d[E/(NBT*beta) + CE]/dtheta + at the relaxed states (free-phase dE/dtheta == 0 exactly). Divergence guard skips bad batches.""" 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) + with torch.no_grad(): + drift = sum(float((a - b).norm()) for a, b in zip(zp, zs_free)) / max( + sum(float(b.norm()) for b in zs_free), 1e-9) + if not math.isfinite(drift) or drift > 0.5: # nudged displacement should be O(beta) + for p in all_params: p.grad = None + return free_ce # skip batch (guard) + prev = tok(x) + pos(torch.arange(args.T, device=dev))[None] + E = 0.0 + for z, b in zip(zp, blocks): E = E + 0.5 * ((z - b(prev, mask)) ** 2).sum(); prev = z + obj = E / (NBT * args.beta) + F.cross_entropy(readout(zp[-1]).reshape(-1, vocab), y.reshape(-1)) + gs = torch.autograd.grad(obj, all_params, allow_unused=True) + for p, g in zip(all_params, gs): + p.grad = g return free_ce @torch.no_grad() @@ -116,7 +131,7 @@ if args.wandb: 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} ' +print(f'[{args.tag}] cascade-EP(EQUILIBRIUM/fb) 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): diff --git a/ep_run/cascade_probe.py b/ep_run/cascade_probe.py index dda4059..d6d9fea 100644 --- a/ep_run/cascade_probe.py +++ b/ep_run/cascade_probe.py @@ -18,7 +18,7 @@ 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', 'zil'], default='jacobi') +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) @@ -26,6 +26,7 @@ ap.add_argument('--ckpt', type=str, default='') # gate at a casc_bp_trai 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('--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' @@ -87,10 +88,35 @@ def local_grad(zs, z0, l, beta, y): 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): @@ -108,13 +134,55 @@ def relax(z0, y, beta): (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) - zs[l] = (zs[l] - args.eta * bufs[l]).detach() + 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): |
