summaryrefslogtreecommitdiff
path: root/experiments
diff options
context:
space:
mode:
Diffstat (limited to 'experiments')
-rw-r--r--experiments/run.py11
-rw-r--r--experiments/smoke.py32
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))