From ddff6bf7d31b0dd28bde1d057939c2c4f9b359ba Mon Sep 17 00:00:00 2001 From: Yuren Hao Date: Thu, 9 Jul 2026 07:44:14 -0500 Subject: =?UTF-8?q?cascade=20equilibrium=20solver=20solved:=20fb=20(forwar?= =?UTF-8?q?d-backward=20message=20passing)=20K=3D3=20beta=3D0.003=20?= =?UTF-8?q?=E2=80=94=20gates=201.0000/0.9990/0.9946=20across=20BP=20trajec?= =?UTF-8?q?tory,=20L12=200.9999;=20trainer=20v2=20single-sided=20EP=20read?= =?UTF-8?q?out=20+=20divergence=20guard=20(~5x=20BP)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Co-Authored-By: Claude Fable 5 Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn --- ep_run/cascade_probe.py | 72 +++++++++++++++++++++++++++++++++++++++++++++++-- 1 file changed, 70 insertions(+), 2 deletions(-) (limited to 'ep_run/cascade_probe.py') 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): -- cgit v1.2.3