"""Disambiguate the parity gap: real bug vs FD-amplified fp noise. (a) single-eval exactness: at a synthetic two-phase state Z=[zs+d, zs-d], compare the full doubled-batch correction (Jv-JTv) against the halved+mirrored one. Exact math => allclose at fp level. (b) estimator noise floor: perturb zs by 1e-6 relative and rerun the ORIGINAL holo_a_track — if a_best moves by ~the same 0.4 rel, the parity gap is the estimator's intrinsic FD sensitivity, not a bug.""" import torch, torch.func as tf import lt_ep_train as L from holo_ep import holo_a_track torch.manual_seed(0) blk = L.EQBlock(512, 16, 256, 256, c=1.0, attn_mode='thick'); blk.qknorm = True ck = torch.load('runs/redx_traj/s2000.pt', map_location=L.dev) with torch.no_grad(): for p, w in zip(blk.allp, ck['allp']): p.copy_(w.to(L.dev)) torch.manual_seed(42) idx, y = L.get_batch('train', 24, 256) xin = blk.embed(idx).detach() zs = L.relax(blk, xin.clone(), xin, 150, 0.1) B = zs.size(0) fnc = lambda zz: blk.nc_force(zz) # (a) single-eval exactness torch.manual_seed(7) d = 0.02 * torch.randn_like(zs) Z = torch.cat([zs + d, zs - d], 0) zbar = 0.5 * (Z[:B] + Z[B:]) zb2 = torch.cat([zbar, zbar], 0) v = (Z - zb2).contiguous() with torch.no_grad(): _, Jv = tf.jvp(fnc, (zb2,), (v,)) JTv = tf.vjp(fnc, zb2)[1](v)[0] corr_full = Jv - JTv v0 = (Z[:B] - zbar).contiguous() _, Jv0 = tf.jvp(fnc, (zbar,), (v0,)) JTv0 = tf.vjp(fnc, zbar)[1](v0)[0] corr_half = torch.cat([Jv0 - JTv0, -(Jv0 - JTv0)], 0) rel_a = ((corr_full - corr_half).norm() / (corr_full.norm() + 1e-12)).item() print(f"(a) single-eval: rel={rel_a:.2e} ({'EXACT — gap is fp/FD noise' if rel_a < 1e-4 else 'REAL BUG'})", flush=True) # (b) estimator noise floor of the ORIGINAL r, T2, eps = 0.02, 40, 0.1 a_ref, _ = holo_a_track(blk, zs, xin, y, r, T2, eps) zs_p = zs + 1e-6 * zs.norm() / (zs.numel() ** 0.5) * torch.randn_like(zs) a_prt, _ = holo_a_track(blk, zs_p, xin, y, r, T2, eps) rel_b = ((a_prt - a_ref).norm() / (a_ref.norm() + 1e-12)).item() cos_b = torch.nn.functional.cosine_similarity(a_ref.flatten(), a_prt.flatten(), dim=0).item() print(f"(b) orig-vs-orig under 1e-6 state noise: rel={rel_b:.2e} cos={cos_b:.6f}", flush=True) print("HOLO_FAST_PROBE_DONE", flush=True)