summaryrefslogtreecommitdiff
path: root/ep_run
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-07-09 09:12:44 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-07-09 09:12:44 -0500
commit7e338ed314d7d73fdfaf32daa5bc105127a42c99 (patch)
tree05bed72cbe27730d8c3beeb9b96ca603689e8a31 /ep_run
parent4aefd494812d3e336cb81765ef75c9dc0eff80af (diff)
casc_eq_train v4: quality-governed estimator — periodic measured cos(EP,BP) drives K/beta (spend when quality drops, relax when abundant); guards reduced to sanity-only
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run')
-rw-r--r--ep_run/casc_eq_train.py51
1 files changed, 28 insertions, 23 deletions
diff --git a/ep_run/casc_eq_train.py b/ep_run/casc_eq_train.py
index 25cfeb6..96b7ab3 100644
--- a/ep_run/casc_eq_train.py
+++ b/ep_run/casc_eq_train.py
@@ -73,13 +73,12 @@ def tok_sigma(iters=8):
v = W.t() @ u; sig = v.norm(); v /= max(sig, 1e-12)
return float(sig)
-def relax(z0, zs_free, y, beta):
- """adaptive-K fb rounds: backward feedback refresh + forward rebuild, until the
- per-round state delta contracts below 5% of round-1 (or kmax). Returns (zs, rounds, contracted)."""
+def relax(z0, zs_free, y, beta, K):
+ """K fb rounds: backward feedback refresh + forward rebuild. State oscillation is
+ HARMLESS for the theta-readout (probe-verified); no contraction verdict."""
zs = [z.clone() for z in zs_free]
d = [None] * args.L
- d1 = None; delta = 0.0
- for k in range(args.kmax):
+ for k in range(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()
@@ -87,22 +86,13 @@ def relax(z0, zs_free, y, beta):
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()
- delta = 0.0
with torch.no_grad():
prev = z0
for l in range(args.L):
rebuilt = blocks[l](prev, mask) + d[l]
- new = (1 - args.geta) * zs[l] + args.geta * rebuilt if args.geta < 1.0 else rebuilt
- delta += float((new - zs[l]).norm())
- zs[l] = new
+ zs[l] = (1 - args.geta) * zs[l] + args.geta * rebuilt if args.geta < 1.0 else rebuilt
prev = zs[l]
- if k == 0:
- d1 = max(delta, 1e-12)
- elif delta < 0.05 * d1 and k + 1 >= args.K:
- return zs, k + 1, True
- elif delta > 3.0 * d1:
- return zs, k + 1, False
- return zs, args.kmax, delta < 0.5 * d1
+ return zs
def dFdtheta(zs, x, y, beta):
"""dF/dtheta at fixed relaxed states (z0 rebuilt WITH graph so emb gets its E-path grad)."""
@@ -114,30 +104,41 @@ def dFdtheta(zs, x, y, beta):
return [g if g is not None else None for g in gs]
SIG0 = None
+GOV = {'K': None, 'bscale': 1.0, 'gema': None}
def ep_step(x, y):
- """single-sided adaptive EP: beta_t = beta0*sig0^2/sig_tok^2 (top-CE stiffness compensation),
- adaptive-K fb relax, grad at the relaxed states. Returns (free_ce, beta_t, rounds, ok)."""
+ """single-sided EP with a QUALITY-GOVERNED estimator: beta_t = beta0*bscale*sig0^2/sig^2,
+ K = GOV['K'] fb rounds; guard = finiteness + drift + grad-norm sanity only."""
global SIG0
+ if GOV['K'] is None: GOV['K'] = args.K
sig = tok_sigma()
if SIG0 is None: SIG0 = sig
- beta_t = args.beta * (SIG0 * SIG0) / max(sig * sig, 1e-9)
+ beta_t = args.beta * GOV['bscale'] * (SIG0 * SIG0) / max(sig * sig, 1e-9)
z0, zs_free = free_states(x)
free_ce = F.cross_entropy(readout(zs_free[-1]).reshape(-1, vocab), y.reshape(-1)).item()
- zp, rounds, ok = relax(z0, zs_free, y, +beta_t)
+ zp = relax(z0, zs_free, y, +beta_t, GOV['K'])
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 ok) or (not math.isfinite(drift)) or drift > 0.5:
+ if (not math.isfinite(drift)) or drift > 0.5:
for p in all_params: p.grad = None
- return free_ce, beta_t, rounds, False
+ return free_ce, beta_t, GOV['K'], False
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 * beta_t) + F.cross_entropy(readout(zp[-1]).reshape(-1, vocab), y.reshape(-1))
gs = torch.autograd.grad(obj, all_params, allow_unused=True)
+ gn = 0.0
+ for g in gs:
+ if g is not None: gn += float((g ** 2).sum())
+ gn = gn ** 0.5
+ if GOV['gema'] is None: GOV['gema'] = gn
+ if not math.isfinite(gn) or gn > 8 * GOV['gema']:
+ for p in all_params: p.grad = None
+ return free_ce, beta_t, GOV['K'], False
+ GOV['gema'] = 0.99 * GOV['gema'] + 0.01 * gn
for p, g in zip(all_params, gs):
p.grad = g
- return free_ce, beta_t, rounds, True
+ return free_ce, beta_t, GOV['K'], True
def bp_gate(x, y):
"""true BP grads for telemetry cos (called before opt.step; reads p.grad separately)."""
@@ -182,6 +183,10 @@ for step in range(args.steps + 1):
if p.grad is None or g is None: continue
num += float((p.grad * g).sum()); den1 += float((p.grad ** 2).sum()); den2 += float((g ** 2).sum())
gcos = num / max((den1 ** 0.5) * (den2 ** 0.5), 1e-12)
+ if gcos < 0.97: # estimator governor: spend more
+ GOV['K'] = min(GOV['K'] + 2, args.kmax); GOV['bscale'] = max(GOV['bscale'] * 0.7, 0.05)
+ elif gcos > 0.995 and GOV['K'] > args.K: # relax back when quality is abundant
+ GOV['K'] -= 1; GOV['bscale'] = min(GOV['bscale'] * 1.05, 1.0)
torch.nn.utils.clip_grad_norm_(all_params, 1.0)
opt.step(); sched.step(); opt.zero_grad(set_to_none=True)
if step % args.log == 0: