summaryrefslogtreecommitdiff
path: root/experiments/baseline_smoke.py
diff options
context:
space:
mode:
Diffstat (limited to 'experiments/baseline_smoke.py')
-rw-r--r--experiments/baseline_smoke.py12
1 files changed, 11 insertions, 1 deletions
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__":