diff options
| author | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 04:31:59 -0500 |
|---|---|---|
| committer | Yuren Hao <yurenh2@illinois.edu> | 2026-07-09 04:31:59 -0500 |
| commit | 57ad5f340a5148417d5f50420af0f7982ddc8105 (patch) | |
| tree | 03e3bbfbf754de87d80ac05c0ccfbc87d490d3a0 /ep_run/cascade_probe.py | |
| parent | cf96818498898575289a9430ad326d9450d3f3a0 (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/cascade_probe.py')
| -rw-r--r-- | ep_run/cascade_probe.py | 57 |
1 files changed, 45 insertions, 12 deletions
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): |
