diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 02:14:31 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-22 02:14:31 -0500 |
| commit | 065c891be816b92a43065bafe89ad2aa266161b1 (patch) | |
| tree | 0364dc4039ca53d8d82e5187ec5ce2dbd32822f5 | |
| parent | 7e8d314c0fb8e82730e24da153e8b73c3af4ecd6 (diff) | |
experiments: isolate apical traffic seeds
| -rw-r--r-- | experiments/run.py | 4 | ||||
| -rw-r--r-- | experiments/smoke.py | 20 |
2 files changed, 23 insertions, 1 deletions
diff --git a/experiments/run.py b/experiments/run.py index aa509e4..a44e72e 100644 --- a/experiments/run.py +++ b/experiments/run.py @@ -101,7 +101,7 @@ def build(args, device): 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, - traffic_mode=args.traffic_mode) + traffic_mode=args.traffic_mode, nuis_seed=args.traffic_seed) if args.mode == "dfa": cfg = dfa_config(eta=args.eta, momentum=args.momentum) elif args.mode == "sdil": @@ -441,6 +441,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_seed", type=int, default=1234, + help="ordinary apical-traffic projection seed; independent of model/data seeds") p.add_argument("--traffic_mode", default="soma", choices=["none", "soma", "topdown", "mixed"]) p.add_argument("--predictor_mode", default="diagonal", choices=["diagonal", "full"]) diff --git a/experiments/smoke.py b/experiments/smoke.py index 83ee896..025d3cc 100644 --- a/experiments/smoke.py +++ b/experiments/smoke.py @@ -113,12 +113,32 @@ def check_topdown_predictor(): assert final[-1] < 0.03 +def check_traffic_seed_isolation(): + """Traffic-family seeds must not change feedforward initialization/data.""" + net_a = SDILNet([20, 32, 32, 32, 5], device="cpu", seed=6, nuis_rho=1.0, + nuis_seed=41, residual=True, traffic_mode="topdown") + net_b = SDILNet([20, 32, 32, 32, 5], device="cpu", seed=6, nuis_rho=1.0, + nuis_seed=42, residual=True, traffic_mode="topdown") + x = torch.randn(64, 20) + for wa, wb in zip(net_a.W, net_b.W): + assert torch.equal(wa, wb) + ha = net_a.forward(x)["h"] + hb = net_b.forward(x)["h"] + for xa, xb in zip(ha, hb): + assert torch.equal(xa, xb) + traffic_a = net_a.apical_traffic(0, ha[1], ha[-2]) + traffic_b = net_b.apical_traffic(0, hb[1], hb[-2]) + assert not torch.equal(traffic_a, traffic_b) + print("CHECK0d traffic seed changes feedback only: passed") + + def main(): torch.manual_seed(0) dev = "cpu" check_residual_local_jacobian() check_neutral_predictor() check_topdown_predictor() + check_traffic_seed_isolation() print("loading MNIST subset...") tr, te, n_in, n_out = get_dataset("mnist", batch_size=128, device=dev) xb, yb = next(iter(tr)) |
