diff options
Diffstat (limited to 'ep_run/lt_ep_train.py')
| -rw-r--r-- | ep_run/lt_ep_train.py | 7 |
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 |
