diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 01:36:37 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 01:36:37 -0500 |
| commit | ee57b4ca69ace14ffadfcfde947112278001feef (patch) | |
| tree | 39ebc1ce1f7c0ba975a02a2b59ca40e11853b593 /experiments/smoke.py | |
| parent | 3504297a39ffef8c167d3458ad71ac1e6146de58 (diff) | |
feat: model endogenous mixed apical traffic
Diffstat (limited to 'experiments/smoke.py')
| -rw-r--r-- | experiments/smoke.py | 32 |
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)) |
