summaryrefslogtreecommitdiff
path: root/ep_run/lt_ep_train.py
diff options
context:
space:
mode:
authorYuren Hao <yurenh2@illinois.edu>2026-07-05 03:53:57 -0500
committerYuren Hao <yurenh2@illinois.edu>2026-07-05 03:53:57 -0500
commitcbecb171b1af77fe6510fd60f1dfa4c7938a20d6 (patch)
treeadbe7b5ad4514930f510267963213afd21b9ebc4 /ep_run/lt_ep_train.py
parent1cea113ef78d2703b024a088703a06ac5d235c5c (diff)
holofast: exact halved-jvp AEP track — 1.55x nudged phase, gradient-gate passed
holo_a_track computed the doubled-batch jvp/vjp on [v0; -v0] at a shared anchor zbar — exact antisymmetric redundancy (phase deviations from the common mode are exact negatives). holo_a_track_fast computes at batch B and mirrors: single-eval parity 6e-7 (exact); trajectory-level 45% divergence SHARED with the original's own FD noise floor (1e-6 state noise -> 49% self-divergence — the 2r=0.04 finite difference amplifies fp noise; training averages it via pema/momentum). Ship gate: cos(EP,BPTT) orig vs fast indistinguishable (0.907/0.912, 0.853/0.853, 0.918/0.918). Timing 6.43->4.16s on the T2=40 nudged phase (contended GPU, relative). --holofast flag, default off; queued ablation arms deliberately stay on orig for fidelity. Co-Authored-By: Claude Fable 5 <noreply@anthropic.com> Claude-Session: https://claude.ai/code/session_014FAPDWQ49M5Ye3NpTndTpn
Diffstat (limited to 'ep_run/lt_ep_train.py')
-rw-r--r--ep_run/lt_ep_train.py7
1 files changed, 5 insertions, 2 deletions
diff --git a/ep_run/lt_ep_train.py b/ep_run/lt_ep_train.py
index d56af1d..e7155d3 100644
--- a/ep_run/lt_ep_train.py
+++ b/ep_run/lt_ep_train.py
@@ -180,12 +180,13 @@ def ep_step(blk, idx, y, T1, T2, eps, beta, jacreg=0.0, holo=0, hr=0.02, t1max=0
z = z + eps * f
return z.detach()
if holo == 2 and t2sel > 0: # adaptive-T2, phase-batched fast path (validated ==)
- from holo_ep import holo_a_select2, holo_a_track
+ from holo_ep import holo_a_select2, holo_a_track, holo_a_track_fast
K = max(1, getattr(blk, 'navg', 1)) # restart-averaging: noise / sqrt(K)
acc = None
for _ in range(K):
if getattr(blk, 'track', False): # common-mode-tracking AEP (loose-tolerant)
- ai, _ = holo_a_track(blk, zs, xin0, y, hr, t2sel, eps)
+ _tr = holo_a_track_fast if getattr(blk, 'holofast', False) else holo_a_track
+ ai, _ = _tr(blk, zs, xin0, y, hr, t2sel, eps) # fast = exact halved-jvp mirror (1.55x)
else:
ai, _ = holo_a_select2(blk, zs, xin0, y, hr, t2sel, eps, li=getattr(blk, 'li_avg', 0))
acc = ai if acc is None else acc + ai
@@ -419,6 +420,7 @@ def main():
ap.add_argument('--li_avg', type=int, default=0) # lock-in integration window (0=snapshot mode)
ap.add_argument('--navg', type=int, default=1) # restart-averaged contrast estimates per update
ap.add_argument('--track', action='store_true') # common-mode-tracking AEP correction
+ ap.add_argument('--holofast', action='store_true') # exact halved-jvp track (1.55x nudged phase; parity = FD noise floor)
ap.add_argument('--rt_final', type=float, default=0.0) # anneal res_target to this (0=off), 25%-75% of run
ap.add_argument('--nudge_brake', type=float, default=0.0) # kappa: anchor spring during nudge (Tikhonov adjoint)
ap.add_argument('--init_ckpt', type=str, default='') # warm-start weights from a saved ckpt
@@ -493,6 +495,7 @@ def main():
blk.li_avg = cfg.li_avg
blk.navg = cfg.navg
blk.track = cfg.track
+ blk.holofast = cfg.holofast
blk.nbrake = cfg.nudge_brake
blk.qknorm = cfg.qknorm
if cfg.resinit != 1.0: # near-identity block at init (contractive) -> stable big-width start