summaryrefslogtreecommitdiff
path: root/ep_run/cascade_probe.py
diff options
context:
space:
mode:
Diffstat (limited to 'ep_run/cascade_probe.py')
-rw-r--r--ep_run/cascade_probe.py57
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):