summaryrefslogtreecommitdiff
path: root/experiments/run.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-21 06:38:55 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-21 06:38:55 -0500
commit215acff518eea2d6d8a5bd7c90c398f49cc3b41b (patch)
tree8c79739d8b0a614ce76e685a3eff323d4cdb1742 /experiments/run.py
chore: capture initial SDIL project state
Diffstat (limited to 'experiments/run.py')
-rw-r--r--experiments/run.py177
1 files changed, 177 insertions, 0 deletions
diff --git a/experiments/run.py b/experiments/run.py
new file mode 100644
index 0000000..153ca41
--- /dev/null
+++ b/experiments/run.py
@@ -0,0 +1,177 @@
+"""
+SDIL main training / diagnostics driver.
+
+Trains one of {bp, dfa, sdil} (sdil with ablation flags) on MNIST/FashionMNIST,
+logging the quantities that actually test the hypothesis:
+ - train loss, test accuracy
+ - per-hidden-layer cos(innovation r_l, -grad h_l) <- the headline metric
+ - cos(raw apical a_l, -grad) and cos(A_l c, -grad) <- residualization ablation
+ - single-step loss-decrease ratio vs exact GD
+
+Everything is JSON-logged for later plotting.
+"""
+import argparse
+import json
+import os
+import sys
+import time
+
+import torch
+
+sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
+from sdil.core import SDILNet, SDILConfig, sdil_step, neutral_p_update
+from sdil.baselines import BPNet, dfa_config, evaluate
+from sdil import probes
+from sdil.data import get_dataset, onehot, make_hierarchical, make_teacher_student
+
+
+def build(args, device):
+ sizes = [args.n_in] + [args.width] * args.depth + [10]
+ if args.mode == "bp":
+ net = BPNet(sizes, act=args.act, device=device, seed=args.seed,
+ w_scale=args.w_scale, nuis_rho=0.0, residual=bool(args.residual))
+ cfg = SDILConfig(eta=args.eta, momentum=args.momentum)
+ return net, cfg
+ 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))
+ if args.mode == "dfa":
+ cfg = dfa_config(eta=args.eta, momentum=args.momentum)
+ elif args.mode == "sdil":
+ cfg = SDILConfig(
+ eta=args.eta, eta_A=args.eta_A, eta_P=args.eta_P,
+ use_residual=bool(args.use_residual), learn_A=bool(args.learn_A),
+ learn_P=bool(args.learn_P), pert_sigma=args.pert_sigma,
+ pert_every=args.pert_every, pert_ndirs=args.pert_ndirs,
+ momentum=args.momentum, settle_steps=args.settle_steps,
+ kappa=args.kappa, feedback=args.feedback,
+ p_update_on_neutral=bool(args.p_neutral),
+ normalize_delta=bool(args.normalize_delta))
+ else:
+ raise ValueError(args.mode)
+ return net, cfg
+
+
+def train(args):
+ device = args.device
+ torch.manual_seed(args.seed)
+ train_loader, test_loader, n_in, n_out = get_dataset(
+ args.dataset, batch_size=args.batch_size, device=device)
+ args.n_in = n_in
+ net, cfg = build(args, device)
+
+ # a fixed probe batch for stable alignment tracking
+ px, py = next(iter(test_loader))
+ px, py = px[:args.probe_bs].to(device), py[:args.probe_bs].to(device)
+ poh = onehot(py, n_out, device=device)
+
+ log = {"args": vars(args), "steps": [], "final": {}}
+ step = 0
+ prev_error = None
+ t0 = time.time()
+
+ # predictor warmup on neutral-period (c=0) drive, so P cancels the apical
+ # nuisance before task plasticity relies on the residual (no-op when rho=0).
+ if args.mode == "sdil" and args.learn_P and args.p_warmup_steps > 0 and args.nuis_rho > 0:
+ it = iter(train_loader)
+ for _ in range(args.p_warmup_steps):
+ try:
+ wx, _ = next(it)
+ except StopIteration:
+ it = iter(train_loader)
+ wx, _ = next(it)
+ neutral_p_update(net, wx.to(device), args.p_warmup_eta)
+ for epoch in range(args.epochs):
+ for x, y in train_loader:
+ x, y = x.to(device), y.to(device)
+ yoh = onehot(y, n_out, device=device)
+ if args.mode == "bp":
+ loss = net.bp_step(x, y, cfg.eta, momentum=cfg.momentum)
+ else:
+ loss, aux = sdil_step(net, x, y, yoh, cfg, step, prev_error=prev_error)
+ prev_error = aux["error"]
+
+ if step % args.log_every == 0:
+ rec = {"step": step, "epoch": epoch, "train_loss": float(loss)}
+ if args.mode != "bp":
+ al = probes.alignment_report(net, px, py, poh, cfg)
+ rec["cos_r_negg"] = al["cos_r_negg"]
+ rec["cos_apical_negg"] = al["cos_apical_negg"]
+ rec["cos_Ac_negg"] = al["cos_Ac_negg"]
+ rec["r_norm"] = al["r_norm"]
+ 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)
+ step += 1
+ if args.max_steps and step >= args.max_steps:
+ break
+ if args.max_steps and step >= args.max_steps:
+ break
+
+ acc, tloss = evaluate(net, test_loader)
+ msg = f"[{args.tag}] epoch {epoch} step {step} loss {loss:.4f} test_acc {acc:.4f}"
+ if args.mode != "bp":
+ al = probes.alignment_report(net, px, py, poh, cfg)
+ meancos = sum(al["cos_r_negg"]) / len(al["cos_r_negg"])
+ msg += f" mean_cos(r,-g) {meancos:+.3f} per-layer {['%.2f'%v for v in al['cos_r_negg']]}"
+ print(msg, flush=True)
+ log["steps"].append({"epoch_end": epoch, "step": step, "test_acc": acc, "test_loss": tloss})
+
+ acc, tloss = evaluate(net, test_loader)
+ log["final"] = {"test_acc": acc, "test_loss": tloss, "wall_s": time.time() - t0}
+ if args.mode != "bp":
+ al = probes.alignment_report(net, px, py, poh, cfg)
+ log["final"]["cos_r_negg"] = al["cos_r_negg"]
+ log["final"]["cos_apical_negg"] = al["cos_apical_negg"]
+ log["final"]["cos_Ac_negg"] = al["cos_Ac_negg"]
+ os.makedirs(args.outdir, exist_ok=True)
+ outpath = os.path.join(args.outdir, f"{args.tag}.json")
+ with open(outpath, "w") as f:
+ json.dump(log, f)
+ print(f"[{args.tag}] DONE test_acc={acc:.4f} -> {outpath}", flush=True)
+ return log
+
+
+def get_args():
+ p = argparse.ArgumentParser()
+ p.add_argument("--mode", default="sdil", choices=["bp", "dfa", "sdil"])
+ p.add_argument("--dataset", default="mnist", choices=["mnist", "fmnist", "cifar10"])
+ p.add_argument("--depth", type=int, default=3) # hidden layers
+ p.add_argument("--width", type=int, default=256)
+ p.add_argument("--act", default="tanh", choices=["tanh", "gelu", "silu"])
+ p.add_argument("--residual", type=int, default=0) # skip connections (deep no-BN)
+ p.add_argument("--epochs", type=int, default=15)
+ p.add_argument("--batch_size", type=int, default=128)
+ p.add_argument("--eta", type=float, default=0.05)
+ p.add_argument("--eta_A", type=float, default=0.02)
+ p.add_argument("--eta_P", type=float, default=0.002)
+ p.add_argument("--momentum", type=float, default=0.9)
+ p.add_argument("--w_scale", type=float, default=1.0)
+ p.add_argument("--a_scale", type=float, default=1.0)
+ p.add_argument("--pert_sigma", type=float, default=1e-2)
+ p.add_argument("--pert_every", type=int, default=4)
+ p.add_argument("--pert_ndirs", type=int, default=4)
+ p.add_argument("--use_residual", type=int, default=1)
+ p.add_argument("--learn_A", type=int, default=1)
+ p.add_argument("--learn_P", type=int, default=1)
+ p.add_argument("--p_neutral", type=int, default=1) # P update on neutral (c=0) drive
+ 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("--normalize_delta", type=int, default=0)
+ p.add_argument("--settle_steps", type=int, default=0)
+ p.add_argument("--kappa", type=float, default=0.0)
+ p.add_argument("--feedback", default="error", choices=["error", "error_deriv"])
+ p.add_argument("--seed", type=int, default=0)
+ p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
+ p.add_argument("--log_every", type=int, default=50)
+ p.add_argument("--max_steps", type=int, default=0) # 0 = no cap (smoke only)
+ p.add_argument("--probe_bs", type=int, default=512)
+ p.add_argument("--outdir", default="results")
+ p.add_argument("--tag", default="sdil_run")
+ return p.parse_args()
+
+
+if __name__ == "__main__":
+ train(get_args())