summaryrefslogtreecommitdiff
path: root/ep_run/lt_ep_train.py
diff options
context:
space:
mode:
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