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 | |
| parent | 3504297a39ffef8c167d3458ad71ac1e6146de58 (diff) | |
feat: model endogenous mixed apical traffic
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/run.py | 11 | ||||
| -rw-r--r-- | experiments/smoke.py | 32 |
2 files changed, 42 insertions, 1 deletions
diff --git a/experiments/run.py b/experiments/run.py index b412b04..cc3e475 100644 --- a/experiments/run.py +++ b/experiments/run.py @@ -65,7 +65,8 @@ def build(args, device): net = SDILNet(sizes, act=args.act, device=device, seed=args.seed, w_scale=args.w_scale, a_scale=args.a_scale, nuis_rho=args.nuis_rho, feedback=args.feedback, - residual=bool(args.residual), predictor_mode=args.predictor_mode) + residual=bool(args.residual), predictor_mode=args.predictor_mode, + traffic_mode=args.traffic_mode) if args.mode == "dfa": cfg = dfa_config(eta=args.eta, momentum=args.momentum) elif args.mode == "sdil": @@ -169,6 +170,9 @@ def train(args): rec["cos_apical_negg"] = al["cos_apical_negg"] rec["cos_Ac_negg"] = al["cos_Ac_negg"] rec["r_norm"] = al["r_norm"] + rec["traffic_norm"] = al["traffic_norm"] + rec["traffic_residual_norm"] = al["traffic_residual_norm"] + rec["traffic_r2"] = al["traffic_r2"] if args.mode == "sdil" and step % (args.log_every * 5) == 0: rec["ldr"] = probes.loss_decrease_ratio(net, px, py, poh, cfg, step) log["steps"].append(rec) @@ -201,6 +205,9 @@ def train(args): log["final"]["cos_innovation_negg"] = al["cos_innovation_negg"] log["final"]["cos_apical_negg"] = al["cos_apical_negg"] log["final"]["cos_Ac_negg"] = al["cos_Ac_negg"] + log["final"]["traffic_norm"] = al["traffic_norm"] + log["final"]["traffic_residual_norm"] = al["traffic_residual_norm"] + log["final"]["traffic_r2"] = al["traffic_r2"] os.makedirs(args.outdir, exist_ok=True) outpath = os.path.join(args.outdir, f"{args.tag}.json") with open(outpath, "w") as f: @@ -252,6 +259,8 @@ def get_args(): p.add_argument("--p_warmup_steps", type=int, default=200) # pre-task neutral P warmup p.add_argument("--p_warmup_eta", type=float, default=0.05) p.add_argument("--nuis_rho", type=float, default=0.0) + p.add_argument("--traffic_mode", default="soma", + choices=["none", "soma", "topdown", "mixed"]) p.add_argument("--predictor_mode", default="diagonal", choices=["diagonal", "full"]) p.add_argument("--normalize_delta", type=int, default=0) p.add_argument("--settle_steps", type=int, default=0) 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)) |
