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.py72
1 files changed, 70 insertions, 2 deletions
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):