summaryrefslogtreecommitdiff
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-07-09 07:44:14 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-07-09 07:44:14 -0500
commitddff6bf7d31b0dd28bde1d057939c2c4f9b359ba (patch)
tree64a9fa748b7fec6225d828694220b9f3cf3cb05a
parentaa20066cbec3546236da55ee111d7265f8e6f0d7 (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
-rw-r--r--ep_run/casc_eq_train.py55
-rw-r--r--ep_run/cascade_probe.py72
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):