diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-21 06:38:55 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-21 06:38:55 -0500 |
| commit | 215acff518eea2d6d8a5bd7c90c398f49cc3b41b (patch) | |
| tree | 8c79739d8b0a614ce76e685a3eff323d4cdb1742 /experiments/run.py | |
chore: capture initial SDIL project state
Diffstat (limited to 'experiments/run.py')
| -rw-r--r-- | experiments/run.py | 177 |
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()) |
