From d6a00aab979a50394cd4d5423f7304bb7819ce25 Mon Sep 17 00:00:00 2001 From: YurenHao0426 Date: Tue, 21 Jul 2026 08:26:40 -0500 Subject: fix: match published baseline protocols --- experiments/baseline_smoke.py | 12 +++++++++++- 1 file changed, 11 insertions(+), 1 deletion(-) (limited to 'experiments/baseline_smoke.py') diff --git a/experiments/baseline_smoke.py b/experiments/baseline_smoke.py index 4bc3d85..24bb2a4 100644 --- a/experiments/baseline_smoke.py +++ b/experiments/baseline_smoke.py @@ -7,6 +7,7 @@ import torch sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from sdil.data import onehot from sdil.local_baselines import FANet, PEPITANet, FFNet, EPNet +from sdil import probes def check_fa_residual_transport(): @@ -28,7 +29,10 @@ def check_fa_residual_transport(): assert all(v > 0 for v in res_changes) assert plain_changes[0] == 0 and plain_changes[1] == 0 assert plain_changes[2] > 0 and plain_changes[3] > 0 + alignment = probes.fa_alignment_report(residual, x, y, yoh)["cos_fa_negg"] + assert len(alignment) == 3 and all(torch.isfinite(torch.tensor(alignment))) print("FA residual identity transport:", res_changes) + print("FA measurable hidden alignment:", alignment) def check_pepita_output_rule(): @@ -93,10 +97,16 @@ def check_ep_energy_dynamics(): actual = net._settle(x, y, beta=beta, s=[v.clone() for v in s], T=1) assert all(torch.allclose(a, b, atol=1e-7) for a, b in zip(actual, manual)) old = [w.clone() for w in net.W] - net.train_step(x, y.argmax(1), y, eta=[0.01, 0.005]) + _, free_state = net.train_step(x, y.argmax(1), y, eta=[0.01, 0.005], + return_free_state=True) assert all(torch.isfinite(w).all() for w in net.W) assert any(not torch.equal(a, b) for a, b in zip(net.W, old)) + assert len(free_state) == net.L + continued = net._settle(x, s=free_state, T=1) + restarted = net._settle(x, T=1) + assert any(not torch.equal(a, b) for a, b in zip(continued, restarted)) print("EP one-step -d(E+beta*C)/ds dynamics: exact") + print("EP free particles persist across presentations") if __name__ == "__main__": -- cgit v1.2.3