summaryrefslogtreecommitdiff
path: root/experiments/smoke.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 01:36:37 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-22 01:36:37 -0500
commitee57b4ca69ace14ffadfcfde947112278001feef (patch)
tree39ebc1ce1f7c0ba975a02a2b59ca40e11853b593 /experiments/smoke.py
parent3504297a39ffef8c167d3458ad71ac1e6146de58 (diff)
feat: model endogenous mixed apical traffic
Diffstat (limited to 'experiments/smoke.py')
-rw-r--r--experiments/smoke.py32
1 files changed, 32 insertions, 0 deletions
diff --git a/experiments/smoke.py b/experiments/smoke.py
index 0bdb0ca..83ee896 100644
--- a/experiments/smoke.py
+++ b/experiments/smoke.py
@@ -82,11 +82,43 @@ def check_neutral_predictor():
assert after < 0.03
+def check_topdown_predictor():
+ """A local soma predictor should remove the predictable part of feedback
+ generated by the network's own high-level contextual state."""
+ torch.manual_seed(13)
+ net = SDILNet([20, 32, 32, 32, 5], device="cpu", seed=6, nuis_rho=1.0,
+ predictor_mode="diagonal", residual=True,
+ traffic_mode="topdown")
+ x = torch.randn(256, 20)
+
+ def unexplained_fraction():
+ with torch.no_grad():
+ h = net.forward(x)["h"]
+ context = h[-2]
+ fractions = []
+ for l in range(net.L - 1):
+ traffic = net.apical_traffic(l, h[l + 1], context)
+ residual = traffic - net.baseline(l, h[l + 1])
+ fractions.append((residual.pow(2).mean() / traffic.pow(2).mean()).item())
+ return fractions
+
+ initial = unexplained_fraction()
+ for _ in range(500):
+ neutral_p_update(net, x, 0.02)
+ final = unexplained_fraction()
+ print("CHECK0c top-down traffic unexplained power:")
+ for l, (before, after) in enumerate(zip(initial, final)):
+ print(f" layer {l}: {before:.4f} -> {after:.4f}")
+ assert after < before * 0.8
+ assert final[-1] < 0.03
+
+
def main():
torch.manual_seed(0)
dev = "cpu"
check_residual_local_jacobian()
check_neutral_predictor()
+ check_topdown_predictor()
print("loading MNIST subset...")
tr, te, n_in, n_out = get_dataset("mnist", batch_size=128, device=dev)
xb, yb = next(iter(tr))