diff options
| author | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-21 08:26:40 -0500 |
|---|---|---|
| committer | YurenHao0426 <Blackhao0426@gmail.com> | 2026-07-21 08:26:40 -0500 |
| commit | d6a00aab979a50394cd4d5423f7304bb7819ce25 (patch) | |
| tree | d2feb5affd443147b57f3b31c30fd28149f4f722 /experiments/baseline_run.py | |
| parent | c8a372cae2612839bdbdafcf9c114f52cb4a6de3 (diff) | |
fix: match published baseline protocols
Diffstat (limited to 'experiments/baseline_run.py')
| -rw-r--r-- | experiments/baseline_run.py | 136 |
1 files changed, 114 insertions, 22 deletions
diff --git a/experiments/baseline_run.py b/experiments/baseline_run.py index a96e24d..24ad6ae 100644 --- a/experiments/baseline_run.py +++ b/experiments/baseline_run.py @@ -19,6 +19,7 @@ import torch.nn.functional as F sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) from sdil.data import get_dataset, onehot from sdil.local_baselines import FANet, PEPITANet, FFNet, EPNet +from sdil import probes METHOD_SOURCES = { @@ -39,7 +40,7 @@ METHOD_SOURCES = { "ep": { "paper": "https://www.frontiersin.org/journals/computational-neuroscience/articles/10.3389/fncom.2017.00024/full", "code": "https://github.com/bscellier/Towards-a-Biologically-Plausible-Backprop", - "protocol": "two-phase energy-gradient dynamics; random beta sign", + "protocol": "two-phase energy-gradient dynamics; persistent free particles; random beta sign", }, } @@ -108,10 +109,54 @@ def ep_learning_rates(depth, base_eta=None): return [base_eta / (4 ** l) for l in range(depth + 1)] if depth in canonical: return canonical[depth] - # The original implementation already required rapidly shrinking rates as - # depth grew. Continue its depth-3 geometric schedule for an explicit, - # logged extrapolation rather than pretending there is a canonical one. - return [0.128 / (4 ** l) for l in range(depth + 1)] + raise ValueError("non-canonical EP depth requires an explicit --eta") + + +def ep_dynamics(depth, beta=None, free_steps=None, nudge_steps=None): + """Original-model dynamics for the published 1/2/3-hidden-layer runs.""" + canonical = { + 1: {"beta": 0.5, "free_steps": 20, "nudge_steps": 4}, + 2: {"beta": 1.0, "free_steps": 150, "nudge_steps": 6}, + 3: {"beta": 1.0, "free_steps": 500, "nudge_steps": 8}, + } + if depth not in canonical and (beta is None or free_steps is None or nudge_steps is None): + raise ValueError( + "non-canonical EP depth requires --ep_beta, --ep_free_steps, and --ep_nudge_steps") + defaults = canonical.get(depth, {}) + return { + "beta": defaults["beta"] if beta is None else beta, + "free_steps": defaults["free_steps"] if free_steps is None else free_steps, + "nudge_steps": defaults["nudge_steps"] if nudge_steps is None else nudge_steps, + "canonical_depth": depth in canonical, + } + + +def pepita_protocol(args): + """Resolve only settings that are actually supported by a cited protocol. + + The original PEPITA paper specifies one hidden layer. Deep PEPITA is very + sensitive to architecture-specific learning rates, so silently reusing the + shallow 0.1 default is worse than requiring an explicit value. + """ + eta = args.eta + if eta is None: + if args.depth != 1: + raise ValueError("deep PEPITA requires an explicit --eta") + if args.dataset == "mnist": + eta = 0.1 + elif args.dataset == "cifar10": + eta = 0.01 + else: + raise ValueError("PEPITA on this dataset requires an explicit --eta") + if args.pepita_decay_epochs is not None: + decay_epochs = [int(v) for v in args.pepita_decay_epochs.split(",") if v.strip()] + elif args.depth == 1 and args.dataset == "mnist": + decay_epochs = [60] + elif args.depth == 1 and args.dataset == "cifar10": + decay_epochs = [60, 90] + else: + decay_epochs = [] + return {"eta": eta, "decay_epochs": decay_epochs} def build(args, n_in, device): @@ -128,30 +173,53 @@ def build(args, n_in, device): return FFNet([n_in] + [args.width] * args.depth, device=device, seed=args.seed, threshold=args.ff_threshold) if args.method == "ep": - return EPNet(sizes, device=device, seed=args.seed, beta=args.ep_beta, - dt=args.ep_dt, T_free=args.ep_free_steps, - T_nudge=args.ep_nudge_steps, random_beta_sign=True) + dyn = ep_dynamics(args.depth, args.ep_beta, args.ep_free_steps, + args.ep_nudge_steps) + return EPNet(sizes, device=device, seed=args.seed, beta=dyn["beta"], + dt=args.ep_dt, T_free=dyn["free_steps"], + T_nudge=dyn["nudge_steps"], random_beta_sign=True) raise ValueError(args.method) def train(args): torch.manual_seed(args.seed) device = args.device + pepita_cfg = pepita_protocol(args) if args.method == "pepita" else None + ep_dyn = (ep_dynamics(args.depth, args.ep_beta, args.ep_free_steps, + args.ep_nudge_steps) if args.method == "ep" else None) + use_persistent_ep = args.method == "ep" and bool(args.ep_persistent) train_loader, test_loader, n_in, n_out = get_dataset( - args.dataset, args.batch_size, device=device) + args.dataset, args.batch_size, device=device, + shuffle_train=not use_persistent_ep, + train_limit=args.train_examples or None) net = build(args, n_in, device) canonical = bool(args.canonical_preprocess) + resolved_protocol = {} + if pepita_cfg is not None: + resolved_protocol.update(pepita_cfg) + if ep_dyn is not None: + resolved_protocol.update(ep_dyn) + resolved_protocol["learning_rates"] = ep_learning_rates(args.depth, args.eta) + resolved_protocol["persistent_particles"] = use_persistent_ep + resolved_protocol["train_examples"] = train_loader.n log = { "args": vars(args), + "resolved_protocol": resolved_protocol, "method_source": METHOD_SOURCES[args.method], "provenance": code_provenance(), "steps": [], "final": {}, } + if args.method == "fa": + px, py = next(iter(test_loader)) + px, py = px[:args.probe_bs], py[:args.probe_bs] + pyoh = onehot(py, n_out, device=device) t0 = time.time() if args.method == "ff": eta = 0.03 if args.eta is None else args.eta + resolved_protocol["eta"] = eta + resolved_protocol["epochs_per_layer"] = args.epochs for layer in range(args.depth): for epoch in range(args.epochs): batches = 0 @@ -168,15 +236,21 @@ def train(args): "local_loss": loss}) print(f"[{args.tag}] layer {layer} test_acc {acc:.4f}", flush=True) else: - eta = args.eta - if eta is None: - eta = 0.05 if args.method == "fa" else (0.1 if args.method == "pepita" else None) + if args.method == "fa": + eta = 0.05 if args.eta is None else args.eta + resolved_protocol["eta"] = eta + resolved_protocol["feedback_scale"] = args.feedback_scale + elif args.method == "pepita": + eta = pepita_cfg["eta"] + else: + eta = args.eta ep_etas = ep_learning_rates(args.depth, eta) if args.method == "ep" else None + ep_states = [None] * len(train_loader) if use_persistent_ep else None for epoch in range(args.epochs): - if args.method == "pepita" and epoch in (60, 90): + if args.method == "pepita" and epoch in pepita_cfg["decay_epochs"]: eta *= 0.1 batches = 0 - for x, y in train_loader: + for batch_index, (x, y) in enumerate(train_loader): x = canonical_input(x, args.dataset, args.method, canonical) yoh = onehot(y, n_out, device=device) if args.method == "fa": @@ -184,21 +258,33 @@ def train(args): elif args.method == "pepita": loss = net.pepita_step(x, y, yoh, eta, args.momentum) else: - loss = net.train_step(x, y, yoh, ep_etas) + if use_persistent_ep: + loss, ep_states[batch_index] = net.train_step( + x, y, yoh, ep_etas, free_state=ep_states[batch_index], + return_free_state=True) + else: + loss = net.train_step(x, y, yoh, ep_etas) batches += 1 if args.max_batches and batches >= args.max_batches: break acc, test_loss = evaluate_method(net, test_loader, args.method, args.dataset, canonical) - log["steps"].append({"epoch_end": epoch, "test_acc": acc, - "test_loss": test_loss, "train_loss": loss}) - print(f"[{args.tag}] epoch {epoch} loss {loss:.4f} test_acc {acc:.4f}", - flush=True) + rec = {"epoch_end": epoch, "test_acc": acc, + "test_loss": test_loss, "train_loss": loss} + msg = f"[{args.tag}] epoch {epoch} loss {loss:.4f} test_acc {acc:.4f}" + if args.method == "fa": + al = probes.fa_alignment_report(net, px, py, pyoh) + rec["cos_fa_negg"] = al["cos_fa_negg"] + msg += f" mean_cos(fa,-g) {sum(al['cos_fa_negg']) / len(al['cos_fa_negg']):+.3f}" + log["steps"].append(rec) + print(msg, flush=True) acc, test_loss = evaluate_method(net, test_loader, args.method, args.dataset, canonical) log["final"] = {"test_acc": acc, "test_loss": test_loss, "wall_s": time.time() - t0} + if args.method == "fa": + log["final"].update(probes.fa_alignment_report(net, px, py, pyoh)) os.makedirs(args.outdir, exist_ok=True) path = os.path.join(args.outdir, f"{args.tag}.json") with open(path, "w") as f: @@ -216,6 +302,8 @@ def get_args(): p.add_argument("--epochs", type=int, default=10, help="FF: epochs per greedy layer; other methods: global epochs") p.add_argument("--batch_size", type=int, default=64) + p.add_argument("--train_examples", type=int, default=0, + help="0 uses the full training split") p.add_argument("--max_batches", type=int, default=0) p.add_argument("--eta", type=float, default=None) p.add_argument("--momentum", type=float, default=0.9) @@ -223,13 +311,17 @@ def get_args(): p.add_argument("--residual", type=int, default=0) p.add_argument("--feedback_scale", type=float, default=1.0) p.add_argument("--pepita_f_scale", type=float, default=0.05) + p.add_argument("--pepita_decay_epochs", default=None, + help="comma-separated; deep PEPITA defaults to no decay") p.add_argument("--keep_prob", type=float, default=0.9) p.add_argument("--ff_threshold", type=float, default=2.0) - p.add_argument("--ep_beta", type=float, default=0.5) + p.add_argument("--ep_beta", type=float, default=None) p.add_argument("--ep_dt", type=float, default=0.5) - p.add_argument("--ep_free_steps", type=int, default=20) - p.add_argument("--ep_nudge_steps", type=int, default=4) + p.add_argument("--ep_free_steps", type=int, default=None) + p.add_argument("--ep_nudge_steps", type=int, default=None) + p.add_argument("--ep_persistent", type=int, default=1) p.add_argument("--canonical_preprocess", type=int, default=1) + p.add_argument("--probe_bs", type=int, default=512) p.add_argument("--seed", type=int, default=0) p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") p.add_argument("--outdir", default="results") |
