diff options
Diffstat (limited to 'ep_run/casc_eq_train.py')
| -rw-r--r-- | ep_run/casc_eq_train.py | 55 |
1 files changed, 35 insertions, 20 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): |
