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 | |
chore: capture initial SDIL project state
Diffstat (limited to 'experiments')
| -rw-r--r-- | experiments/analyze.py | 191 | ||||
| -rw-r--r-- | experiments/analyze_depth.py | 134 | ||||
| -rw-r--r-- | experiments/calib_tent.py | 18 | ||||
| -rw-r--r-- | experiments/deep_sweep.sh | 37 | ||||
| -rw-r--r-- | experiments/depth_sweep.sh | 39 | ||||
| -rw-r--r-- | experiments/diag_norm.sh | 15 | ||||
| -rw-r--r-- | experiments/probe_credit.py | 114 | ||||
| -rw-r--r-- | experiments/run.py | 177 | ||||
| -rw-r--r-- | experiments/run_battery.sh | 88 | ||||
| -rw-r--r-- | experiments/run_v2.sh | 19 | ||||
| -rw-r--r-- | experiments/seeds5.sh | 4 | ||||
| -rw-r--r-- | experiments/sigfix.sh | 8 | ||||
| -rw-r--r-- | experiments/smoke.py | 86 | ||||
| -rw-r--r-- | experiments/zoo.py | 83 | ||||
| -rw-r--r-- | experiments/zoo_cpu.sh | 6 |
15 files changed, 1019 insertions, 0 deletions
diff --git a/experiments/analyze.py b/experiments/analyze.py new file mode 100644 index 0000000..67b6f04 --- /dev/null +++ b/experiments/analyze.py @@ -0,0 +1,191 @@ +"""Aggregate results/*.json into summary tables (+ figures if matplotlib is +present). Run on any box with access to the shared NFS home. +Usage: python experiments/analyze.py [results_dir] +""" +import glob +import json +import os +import sys +from collections import defaultdict + +RES = sys.argv[1] if len(sys.argv) > 1 else "results" + +try: + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + HAVE_MPL = True +except Exception: + HAVE_MPL = False + + +def load_all(res): + runs = [] + for p in sorted(glob.glob(os.path.join(res, "*.json"))): + try: + with open(p) as f: + d = json.load(f) + d["_path"] = p + d["_tag"] = os.path.basename(p)[:-5] + runs.append(d) + except Exception as e: + print("skip", p, e) + return runs + + +def mean_final_cos(run): + f = run.get("final", {}) + c = f.get("cos_r_negg") + return sum(c) / len(c) if c else None + + +def group(runs, keyfn): + g = defaultdict(list) + for r in runs: + k = keyfn(r) + if k is not None: + g[k].append(r) + return g + + +def agg(vals): + vals = [v for v in vals if v is not None] + if not vals: + return (None, None, 0) + m = sum(vals) / len(vals) + sd = (sum((v - m) ** 2 for v in vals) / len(vals)) ** 0.5 + return (m, sd, len(vals)) + + +def fmt(m, sd, n): + if m is None: + return " n/a " + return f"{m:.3f}±{sd:.3f}(n{n})" + + +def table_main(runs): + print("\n================ MAIN COMPARISON (mnist depth3 width256) ================") + sel = [r for r in runs if r["args"].get("depth") == 3 and r["args"].get("width") == 256 + and r["args"].get("dataset") == "mnist" and r["args"].get("nuis_rho", 0) == 0 + and r["_tag"].startswith(("main_", "A_main_"))] + by_mode = group(sel, lambda r: r["args"]["mode"]) + print(f"{'mode':6s} {'test_acc':22s} {'mean_cos(r,-g)':22s}") + for mode in ["bp", "dfa", "sdil"]: + rs = by_mode.get(mode, []) + acc = agg([r["final"].get("test_acc") for r in rs]) + cos = agg([mean_final_cos(r) for r in rs]) + print(f"{mode:6s} {fmt(*acc):22s} {fmt(*cos):22s}") + + +def table_nuisance(runs): + print("\n================ RESIDUALIZATION UNDER NUISANCE ================") + print("(sdil; does subtracting the soma-predictable baseline preserve alignment as rho grows?)") + sel = [r for r in runs if r["_tag"].startswith(("nuis_", "A_nuis_"))] + by = group(sel, lambda r: (r["args"].get("nuis_rho"), r["args"].get("use_residual"))) + rhos = sorted({k[0] for k in by}) + print(f"{'rho':>5s} | {'residual=1 acc':22s} {'cos':16s} | {'residual=0 acc':22s} {'cos':16s}") + for rho in rhos: + r1 = by.get((rho, 1), []) + r0 = by.get((rho, 0), []) + a1 = agg([r["final"].get("test_acc") for r in r1]); c1 = agg([mean_final_cos(r) for r in r1]) + a0 = agg([r["final"].get("test_acc") for r in r0]); c0 = agg([mean_final_cos(r) for r in r0]) + print(f"{rho:5.1f} | {fmt(*a1):22s} {fmt(*c1):16s} | {fmt(*a0):22s} {fmt(*c0):16s}") + + +def table_trainedA(runs): + print("\n================ TRAINED vs FIXED APICAL PATHWAY ================") + fixed = [r for r in runs if "fixedA" in r["_tag"]] + af = agg([r["final"].get("test_acc") for r in fixed]); cf = agg([mean_final_cos(r) for r in fixed]) + print(f"{'fixed-A (DFA-like)':22s} acc {fmt(*af):22s} cos {fmt(*cf):16s}") + tr = [r for r in runs if "trainedA_nd" in r["_tag"]] + by = group(tr, lambda r: r["args"].get("pert_ndirs")) + for nd in sorted(by): + rs = by[nd] + a = agg([r["final"].get("test_acc") for r in rs]); c = agg([mean_final_cos(r) for r in rs]) + print(f"trained-A ndirs={nd:<3d} acc {fmt(*a):22s} cos {fmt(*c):16s}") + + +def table_depth(runs): + print("\n================ DEPTH SCALING ================") + sel = [r for r in runs if r["_tag"].startswith("depth_")] + if not sel: + return + by = group(sel, lambda r: (r["args"]["mode"], r["args"]["depth"])) + depths = sorted({k[1] for k in by}) + print(f"{'depth':>6s} | " + " | ".join(f"{m:>26s}" for m in ["bp", "dfa", "sdil"])) + for d in depths: + cells = [] + for m in ["bp", "dfa", "sdil"]: + rs = by.get((m, d), []) + a = agg([r["final"].get("test_acc") for r in rs]) + c = agg([mean_final_cos(r) for r in rs]) + cells.append(f"a{fmt(*a)[:5]} c{('%.2f'%c[0]) if c[0] is not None else ' n/a'}") + print(f"{d:6d} | " + " | ".join(f"{c:>26s}" for c in cells)) + + +def plots(runs): + if not HAVE_MPL: + print("\n[matplotlib not available -> skipping figures]") + return + figdir = os.path.join(RES, "figs") + os.makedirs(figdir, exist_ok=True) + + # Fig 1: alignment trajectory sdil vs dfa (seed0 main) + plt.figure(figsize=(6, 4)) + for tag, lab in [("A_main_sdil_s0", "SDIL (trained A)"), ("A_main_dfa_s0", "DFA (fixed A)"), + ("main_sdil_mnist_s0", "SDIL"), ("main_dfa_mnist_s0", "DFA")]: + r = next((x for x in runs if x["_tag"] == tag), None) + if not r: + continue + xs = [s["step"] for s in r["steps"] if "cos_r_negg" in s] + ys = [sum(s["cos_r_negg"]) / len(s["cos_r_negg"]) for s in r["steps"] if "cos_r_negg" in s] + if xs: + plt.plot(xs, ys, label=lab) + plt.xlabel("step"); plt.ylabel("mean cos(r, -grad)"); plt.title("Gradient alignment over training") + plt.legend(); plt.grid(alpha=.3); plt.tight_layout() + plt.savefig(os.path.join(figdir, "alignment_traj.png"), dpi=120); plt.close() + + # Fig 2: acc & alignment vs rho (residual vs not) + sel = [r for r in runs if r["_tag"].startswith(("nuis_", "A_nuis_"))] + if sel: + by = group(sel, lambda r: (r["args"].get("nuis_rho"), r["args"].get("use_residual"))) + rhos = sorted({k[0] for k in by}) + fig, ax = plt.subplots(1, 2, figsize=(10, 4)) + for ur, lab in [(1, "residual (SDIL)"), (0, "raw apical")]: + accs = [agg([r["final"].get("test_acc") for r in by.get((rho, ur), [])])[0] for rho in rhos] + coss = [agg([mean_final_cos(r) for r in by.get((rho, ur), [])])[0] for rho in rhos] + ax[0].plot(rhos, accs, marker="o", label=lab) + ax[1].plot(rhos, coss, marker="o", label=lab) + ax[0].set_xlabel("nuisance rho"); ax[0].set_ylabel("test acc"); ax[0].set_title("Accuracy vs nuisance") + ax[1].set_xlabel("nuisance rho"); ax[1].set_ylabel("mean cos(r,-grad)"); ax[1].set_title("Alignment vs nuisance") + for a in ax: + a.legend(); a.grid(alpha=.3) + plt.tight_layout(); plt.savefig(os.path.join(figdir, "nuisance.png"), dpi=120); plt.close() + + # Fig 3: alignment vs depth + sel = [r for r in runs if r["_tag"].startswith("depth_")] + if sel: + by = group(sel, lambda r: (r["args"]["mode"], r["args"]["depth"])) + depths = sorted({k[1] for k in by}) + plt.figure(figsize=(6, 4)) + for m in ["dfa", "sdil"]: + ys = [agg([mean_final_cos(r) for r in by.get((m, d), [])])[0] for d in depths] + plt.plot(depths, ys, marker="o", label=m) + plt.xlabel("hidden depth"); plt.ylabel("mean cos(r,-grad)"); plt.title("Alignment vs depth") + plt.legend(); plt.grid(alpha=.3); plt.tight_layout() + plt.savefig(os.path.join(figdir, "depth.png"), dpi=120); plt.close() + print(f"\n[figures written to {figdir}]") + + +def main(): + runs = load_all(RES) + print(f"loaded {len(runs)} runs from {RES}") + table_main(runs) + table_nuisance(runs) + table_trainedA(runs) + table_depth(runs) + plots(runs) + + +if __name__ == "__main__": + main() diff --git a/experiments/analyze_depth.py b/experiments/analyze_depth.py new file mode 100644 index 0000000..c3ad7b1 --- /dev/null +++ b/experiments/analyze_depth.py @@ -0,0 +1,134 @@ +"""Depth / no-BN analysis: accuracy vs depth, and the per-layer 'depth utility' +profile (does the teaching signal reach early layers?). DFA is expected to give +near-random credit (~0.07) at every layer; SDIL should stay aligned, especially +in the hard early layers. +Usage: python experiments/analyze_depth.py <prefix> e.g. depth or deepres +""" +import glob +import json +import os +import sys +from collections import defaultdict + +PFX = sys.argv[1] if len(sys.argv) > 1 else "depth" +RES = "results" + +try: + import matplotlib + matplotlib.use("Agg") + import matplotlib.pyplot as plt + HAVE_MPL = True +except Exception: + HAVE_MPL = False + + +def load(): + runs = {} + for p in sorted(glob.glob(os.path.join(RES, f"{PFX}_*.json"))): + with open(p) as f: + d = json.load(f) + runs[os.path.basename(p)[:-5]] = d + return runs + + +def agg(vals): + vals = [v for v in vals if v is not None] + if not vals: + return None, None, 0 + m = sum(vals) / len(vals) + sd = (sum((v - m) ** 2 for v in vals) / len(vals)) ** 0.5 + return m, sd, len(vals) + + +def main(): + runs = load() + if not runs: + print(f"no runs with prefix {PFX}_") + return + # collect datasets, depths, seeds + metas = [] + for tag, d in runs.items(): + a = d["args"] + metas.append((a["dataset"], a["mode"], a["depth"], a.get("seed", 0), tag)) + datasets = sorted({m[0] for m in metas}) + depths = sorted({m[2] for m in metas}) + + for ds in datasets: + print(f"\n############ {PFX} dataset={ds} ############") + print("\n--- TEST ACCURACY vs DEPTH ---") + print(f"{'depth':>6s} | {'bp':>18s} | {'dfa':>18s} | {'sdil':>18s} | {'dfa->sdil gap':>14s}") + for depth in depths: + cells = {} + for mode in ["bp", "dfa", "sdil"]: + accs = [runs[t]["final"].get("test_acc") for (D, M, DP, S, t) in metas + if D == ds and M == mode and DP == depth] + cells[mode] = agg(accs) + def cell(m): + mm, sd, n = cells[m] + return f"{mm:.4f}±{sd:.3f}" if mm is not None else " n/a " + gap = "" + if cells["dfa"][0] is not None and cells["sdil"][0] is not None: + gap = f"{cells['sdil'][0] - cells['dfa'][0]:+.4f}" + print(f"{depth:6d} | {cell('bp'):>18s} | {cell('dfa'):>18s} | {cell('sdil'):>18s} | {gap:>14s}") + + # depth utility: per-layer alignment averaged over seeds, DFA vs SDIL + print("\n--- DEPTH UTILITY: per-layer cos(r,-g) [input-side ... output-side] ---") + for depth in depths: + for mode in ["dfa", "sdil"]: + perlayer_runs = [runs[t]["final"].get("cos_r_negg") for (D, M, DP, S, t) in metas + if D == ds and M == mode and DP == depth and runs[t]["final"].get("cos_r_negg")] + if not perlayer_runs: + continue + nL = len(perlayer_runs[0]) + mean_pl = [sum(r[i] for r in perlayer_runs) / len(perlayer_runs) for i in range(nL)] + early = sum(mean_pl[:max(1, nL // 3)]) / max(1, nL // 3) + print(f" d{depth:<2d} {mode:4s} early3rd={early:+.3f} | " + " ".join(f"{v:+.2f}" for v in mean_pl)) + + if HAVE_MPL: + figdir = os.path.join(RES, "figs") + os.makedirs(figdir, exist_ok=True) + for ds in datasets: + plt.figure(figsize=(6, 4)) + for mode, style in [("bp", "k-o"), ("dfa", "r-s"), ("sdil", "b-^")]: + ys, es = [], [] + for depth in depths: + accs = [runs[t]["final"].get("test_acc") for (D, M, DP, S, t) in metas + if D == ds and M == mode and DP == depth] + m, sd, n = agg(accs) + ys.append(m); es.append(sd or 0) + if any(v is not None for v in ys): + xs = [d for d, y in zip(depths, ys) if y is not None] + yv = [y for y in ys if y is not None] + ev = [e for y, e in zip(ys, es) if y is not None] + plt.errorbar(xs, yv, yerr=ev, fmt=style, label=mode, capsize=3) + plt.xlabel("depth (hidden layers)"); plt.ylabel("test acc") + plt.title(f"{PFX} {ds}: accuracy vs depth (no-BN)") + plt.legend(); plt.grid(alpha=.3); plt.tight_layout() + fp = os.path.join(figdir, f"{PFX}_{ds}_acc_vs_depth.png") + plt.savefig(fp, dpi=120); plt.close() + print(f"[fig] {fp}") + + # depth-utility figure: early-layer alignment vs depth (the collapse) + plt.figure(figsize=(6, 4)) + for mode, style in [("dfa", "r-s"), ("sdil", "b-^")]: + xs, ys = [], [] + for depth in depths: + pl = [runs[t]["final"].get("cos_r_negg") for (D, M, DP, S, t) in metas + if D == ds and M == mode and DP == depth and runs[t]["final"].get("cos_r_negg")] + if not pl: + continue + nL = len(pl[0]); k = max(1, nL // 3) + early = [sum(r[:k]) / k for r in pl] + xs.append(depth); ys.append(sum(early) / len(early)) + if xs: + plt.plot(xs, ys, style, label=mode, markersize=7) + plt.xlabel("depth (hidden layers)"); plt.ylabel("early-layer cos(r, -grad)") + plt.title(f"{PFX} {ds}: credit reaching early layers vs depth") + plt.ylim(0, 1); plt.legend(); plt.grid(alpha=.3); plt.tight_layout() + fp2 = os.path.join(figdir, f"{PFX}_{ds}_depth_utility.png") + plt.savefig(fp2, dpi=120); plt.close() + print(f"[fig] {fp2}") + + +if __name__ == "__main__": + main() diff --git a/experiments/calib_tent.py b/experiments/calib_tent.py new file mode 100644 index 0000000..a09fe93 --- /dev/null +++ b/experiments/calib_tent.py @@ -0,0 +1,18 @@ +import torch, collections, sys +from sdil.data import make_tentmap, onehot +from sdil.core import SDILNet, SDILConfig, sdil_step +from sdil.baselines import BPNet, dfa_config, evaluate +dev='cpu' +L=int(sys.argv[1]) if len(sys.argv)>1 else 7 +tr,te,ni,no=make_tentmap(levels=L,n_in=3,n_train=50000,n_test=8000,seed=0,batch_size=128,device=dev) +cnt=collections.Counter(te.y.tolist()); print(f"tent levels={L} balance {dict(cnt)}",flush=True) +for W in [6,10,16]: + print(f"--- width={W} PLAIN relu (residual=0) ---",flush=True) + for d in [1,2,3,4,6,8]: + netb=BPNet([ni]+[W]*d+[2],act='relu',device=dev,seed=1,residual=False) + netd=SDILNet([ni]+[W]*d+[2],act='relu',device=dev,seed=1,residual=False) + cfg=dfa_config(eta=0.02,momentum=0.9); s=0 + for e in range(20): + for x,y in tr: + netb.bp_step(x,y,0.02,0.9); sdil_step(netd,x,y,onehot(y,2),cfg,s); s+=1 + print(f" depth={d:2d}: BP {evaluate(netb,te)[0]:.3f} DFA {evaluate(netd,te)[0]:.3f}",flush=True) diff --git a/experiments/deep_sweep.sh b/experiments/deep_sweep.sh new file mode 100644 index 0000000..15f41a4 --- /dev/null +++ b/experiments/deep_sweep.sh @@ -0,0 +1,37 @@ +#!/usr/bin/env bash +# Deep, no-BN (residual) sweep to expose DFA failure vs SDIL. +# Usage: bash deep_sweep.sh <dataset> "<depths>" <residual> <act> "<seeds>" <epochs> <tag> <gpu> +set -u +cd "$(dirname "$0")/.." +PY=/home/yurenh2/miniconda3/envs/ep_pascal/bin/python3 +export OMP_NUM_THREADS=2 +DS="${1:-cifar10}" +DEPTHS="${2:-5 10 20 30}" +RES="${3:-1}" +ACT="${4:-tanh}" +SEEDS="${5:-0}" +EP="${6:-25}" +PFX="${7:-deep}" +OUT=results +LOGS=logs/deep +mkdir -p "$OUT" "$LOGS" +PROG="$LOGS/progress.txt" +echo "=== deep sweep start $(date) ds=$DS depths=[$DEPTHS] res=$RES act=$ACT seeds=[$SEEDS] ep=$EP ===" >> "$PROG" + +run () { + local tag="$1"; shift + if [ -f "$OUT/$tag.json" ]; then echo "skip $tag" >> "$PROG"; return; fi + echo ">>> $tag $(date +%H:%M:%S)" >> "$PROG" + $PY experiments/run.py --tag "$tag" --outdir "$OUT" "$@" > "$LOGS/$tag.log" 2>&1 \ + && grep -h DONE "$LOGS/$tag.log" >> "$PROG" || echo "FAIL $tag" >> "$PROG" +} + +for s in $SEEDS; do + for d in $DEPTHS; do + for m in bp dfa sdil; do + run "${PFX}_${DS}_${m}_d${d}_s${s}" --mode $m --dataset $DS --depth $d --width 256 \ + --act $ACT --residual $RES --epochs $EP --seed $s --pert_ndirs 8 --eta 0.03 + done + done +done +echo "=== deep sweep done $(date) ===" >> "$PROG" diff --git a/experiments/depth_sweep.sh b/experiments/depth_sweep.sh new file mode 100644 index 0000000..f7c4112 --- /dev/null +++ b/experiments/depth_sweep.sh @@ -0,0 +1,39 @@ +#!/usr/bin/env bash +# Locate the regime where DFA fails and SDIL survives: deep, UNNORMALISED tanh +# MLPs. DFA's random feedback loses gradient alignment with depth; SDIL learns +# the apical pathway per-layer by node perturbation, so alignment should hold. +# Usage: bash depth_sweep.sh "<datasets>" "<depths>" "<seeds>" <epochs> <tag_prefix> +set -u +cd "$(dirname "$0")/.." +PY=/home/yurenh2/miniconda3/envs/ep_pascal/bin/python3 +export OMP_NUM_THREADS=2 +DATASETS="${1:-mnist fmnist}" +DEPTHS="${2:-3 5 7 10}" +SEEDS="${3:-0}" +EP="${4:-20}" +PFX="${5:-depth}" +OUT=results +LOGS=logs/depth +mkdir -p "$OUT" "$LOGS" +PROG="$LOGS/progress.txt" +echo "=== depth sweep start $(date) datasets=[$DATASETS] depths=[$DEPTHS] seeds=[$SEEDS] ===" >> "$PROG" + +run () { + local tag="$1"; shift + if [ -f "$OUT/$tag.json" ]; then echo "skip $tag" >> "$PROG"; return; fi + echo ">>> $tag $(date +%H:%M:%S)" >> "$PROG" + $PY experiments/run.py --tag "$tag" --outdir "$OUT" "$@" > "$LOGS/$tag.log" 2>&1 \ + && grep -h DONE "$LOGS/$tag.log" >> "$PROG" || echo "FAIL $tag" >> "$PROG" +} + +for ds in $DATASETS; do + for s in $SEEDS; do + for d in $DEPTHS; do + for m in bp dfa sdil; do + run "${PFX}_${ds}_${m}_d${d}_s${s}" --mode $m --dataset $ds --depth $d \ + --width 256 --epochs $EP --seed $s --pert_ndirs 8 + done + done + done +done +echo "=== depth sweep done $(date) ===" >> "$PROG" diff --git a/experiments/diag_norm.sh b/experiments/diag_norm.sh new file mode 100644 index 0000000..530cf44 --- /dev/null +++ b/experiments/diag_norm.sh @@ -0,0 +1,15 @@ +#!/usr/bin/env bash +cd /home/yurenh2/sdil +mkdir -p /tmp/diag logs/diag +SUM=logs/diag/summary.txt; : > "$SUM" +cfg() { # tag, extra-args... + local tag="$1"; shift + python3 experiments/run.py --device cpu --dataset fmnist --depth 5 --width 128 \ + --epochs 4 --tag "$tag" --outdir /tmp/diag "$@" > logs/diag/$tag.log 2>&1 + grep -h DONE logs/diag/$tag.log >> "$SUM" +} +cfg d_dfa --mode dfa --eta 0.05 +cfg d_sdil_raw --mode sdil --eta 0.05 --normalize_delta 0 --pert_ndirs 8 +cfg d_sdil_n01 --mode sdil --eta 0.1 --normalize_delta 1 --pert_ndirs 8 +cfg d_sdil_n03 --mode sdil --eta 0.3 --normalize_delta 1 --pert_ndirs 8 +echo "ALL DIAG DONE" >> "$SUM" diff --git a/experiments/probe_credit.py b/experiments/probe_credit.py new file mode 100644 index 0000000..f7c553f --- /dev/null +++ b/experiments/probe_credit.py @@ -0,0 +1,114 @@ +""" +Credit-assignment quality via linear probing. Train a net with {bp,dfa,sdil}, +then FREEZE it and fit a closed-form ridge linear classifier on each hidden +layer's representation. If a rule assigns credit to hidden layers well, those +layers' features become linearly separable. DFA gives early layers near-random +teaching (~0.07 cos), so its early/mid features should probe worse than SDIL's, +independent of the final end-to-end accuracy. + +Usage: python experiments/probe_credit.py --dataset fmnist --depth 6 --epochs 15 +""" +import argparse +import json +import os +import sys +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from sdil.core import SDILNet, SDILConfig, sdil_step +from sdil.baselines import BPNet, dfa_config, evaluate +from sdil.data import get_dataset, onehot + + +def ridge_probe(feat_tr, y_tr, feat_te, y_te, n_classes=10, lam=1e-1): + """Closed-form ridge classifier on frozen features -> test accuracy.""" + X = torch.cat([feat_tr, torch.ones(feat_tr.shape[0], 1, device=feat_tr.device)], 1) + Y = onehot(y_tr, n_classes, device=feat_tr.device) + d = X.shape[1] + A = X.t() @ X + lam * torch.eye(d, device=X.device) + W = torch.linalg.solve(A, X.t() @ Y) + Xte = torch.cat([feat_te, torch.ones(feat_te.shape[0], 1, device=feat_te.device)], 1) + pred = (Xte @ W).argmax(1) + return (pred == y_te).float().mean().item() + + +@torch.no_grad() +def all_features(net, x): + """Return list of hidden activations h[1..L-1] for input x (batched).""" + outs = [[] for _ in range(net.L - 1)] + for i in range(0, x.shape[0], 2000): + fwd = net.forward(x[i:i + 2000]) + for l in range(net.L - 1): + outs[l].append(fwd["h"][l + 1]) + return [torch.cat(o) for o in outs] + + +def train_net(args, device): + train_loader, test_loader, n_in, n_out = get_dataset(args.dataset, args.batch_size, device=device) + sizes = [n_in] + [args.width] * args.depth + [10] + if args.mode == "bp": + net = BPNet(sizes, act=args.act, device=device, seed=args.seed, residual=bool(args.residual)) + cfg = SDILConfig(eta=args.eta, momentum=0.9) + else: + net = SDILNet(sizes, act=args.act, device=device, seed=args.seed, residual=bool(args.residual)) + cfg = (dfa_config(eta=args.eta, momentum=0.9) if args.mode == "dfa" + else SDILConfig(eta=args.eta, eta_A=args.eta_A, momentum=0.9, + pert_ndirs=args.pert_ndirs, pert_every=4, + normalize_delta=bool(args.normalize_delta))) + step = 0 + for ep in range(args.epochs): + for x, y in train_loader: + yoh = onehot(y, 10, device=device) + if args.mode == "bp": + net.bp_step(x, y, cfg.eta, momentum=cfg.momentum) + else: + sdil_step(net, x, y, yoh, cfg, step) + step += 1 + acc, _ = evaluate(net, test_loader) + return net, train_loader, test_loader, acc, n_out + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--dataset", default="fmnist") + p.add_argument("--depth", type=int, default=6) + p.add_argument("--width", type=int, default=256) + p.add_argument("--act", default="tanh") + p.add_argument("--residual", type=int, default=0) + 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("--pert_ndirs", type=int, default=8) + p.add_argument("--normalize_delta", type=int, default=1) + p.add_argument("--seed", type=int, default=0) + p.add_argument("--modes", default="bp,dfa,sdil") + p.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu") + p.add_argument("--outdir", default="results") + p.add_argument("--tag", default="probe") + args = p.parse_args() + device = args.device + + results = {"args": vars(args), "probes": {}} + # shared probe data + _, test_loader, _, _ = get_dataset(args.dataset, args.batch_size, device=device) + for mode in args.modes.split(","): + args.mode = mode + net, train_loader, test_loader, acc, n_out = train_net(args, device) + # gather frozen features (subset of train for speed) + xtr = train_loader.x[:10000]; ytr = train_loader.y[:10000] + xte = test_loader.x; yte = test_loader.y + ftr = all_features(net, xtr); fte = all_features(net, xte) + probe_acc = [ridge_probe(ftr[l], ytr, fte[l], yte, n_out) for l in range(net.L - 1)] + results["probes"][mode] = {"end2end_acc": acc, "layer_probe_acc": probe_acc} + print(f"[{mode}] end2end={acc:.4f} per-layer linear-probe acc [in..out]: " + + " ".join(f"{v:.3f}" for v in probe_acc), flush=True) + + os.makedirs(args.outdir, exist_ok=True) + with open(os.path.join(args.outdir, f"{args.tag}.json"), "w") as f: + json.dump(results, f) + print(f"saved -> {args.outdir}/{args.tag}.json", flush=True) + + +if __name__ == "__main__": + main() 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()) diff --git a/experiments/run_battery.sh b/experiments/run_battery.sh new file mode 100644 index 0000000..847d610 --- /dev/null +++ b/experiments/run_battery.sh @@ -0,0 +1,88 @@ +#!/usr/bin/env bash +# SDIL experiment battery. Runs sequentially on ONE GPU (nets are tiny). +# Usage: bash run_battery.sh <wave> wave in {A,B,all} +set -u +cd "$(dirname "$0")/.." # -> /home/yurenh2/sdil +PY=/home/yurenh2/miniconda3/envs/ep_pascal/bin/python3 +export OMP_NUM_THREADS=2 +WAVE="${1:-all}" +OUT=results +LOGS=logs/battery +mkdir -p "$OUT" "$LOGS" +SEEDS="0 1 2" +EP=20 +PROG="$LOGS/progress.txt" +echo "=== battery start wave=$WAVE $(date) ===" >> "$PROG" + +run () { # run <tag> <args...> + local tag="$1"; shift + if [ -f "$OUT/$tag.json" ]; then echo "skip $tag (exists)" >> "$PROG"; return; fi + echo ">>> $tag $(date +%H:%M:%S)" >> "$PROG" + $PY experiments/run.py --tag "$tag" --outdir "$OUT" "$@" > "$LOGS/$tag.log" 2>&1 \ + && grep -h "DONE" "$LOGS/$tag.log" >> "$PROG" \ + || echo "FAIL $tag" >> "$PROG" +} + +# ---------------- Wave A: core claims, seed 0 (fast health + headline) -------- +if [ "$WAVE" = "A" ] || [ "$WAVE" = "all" ]; then + for m in bp dfa sdil; do + run "A_main_${m}_s0" --mode $m --dataset mnist --depth 3 --width 256 --epochs $EP --seed 0 + done + # residualization under nuisance (the Harnett-specific claim), seed 0 + for rho in 0 2 8; do + for ur in 0 1; do + run "A_nuis_rho${rho}_res${ur}_s0" --mode sdil --nuis_rho $rho --use_residual $ur \ + --depth 3 --width 256 --epochs $EP --seed 0 + done + done + # trained vs fixed apical pathway, seed 0 + run "A_fixedA_s0" --mode sdil --learn_A 0 --depth 3 --width 256 --epochs $EP --seed 0 + for nd in 1 4 16; do + run "A_trainedA_nd${nd}_s0" --mode sdil --learn_A 1 --pert_ndirs $nd --depth 3 --width 256 --epochs $EP --seed 0 + done +fi + +# ---------------- Wave B: full grid, all seeds ------------------------------ +if [ "$WAVE" = "B" ] || [ "$WAVE" = "all" ]; then + for s in $SEEDS; do + # E1 main comparison + fashion-mnist + for m in bp dfa sdil; do + run "main_${m}_mnist_s${s}" --mode $m --dataset mnist --depth 3 --width 256 --epochs $EP --seed $s + run "main_${m}_fmnist_s${s}" --mode $m --dataset fmnist --depth 3 --width 256 --epochs $EP --seed $s + done + # E2 residualization x nuisance + for rho in 0 1 2 4 8; do + for ur in 0 1; do + run "nuis_rho${rho}_res${ur}_s${s}" --mode sdil --nuis_rho $rho --use_residual $ur \ + --depth 3 --width 256 --epochs $EP --seed $s + done + done + # E3 trained vs fixed A, perturbation directions + run "abl_fixedA_s${s}" --mode sdil --learn_A 0 --depth 3 --width 256 --epochs $EP --seed $s + for nd in 1 4 16; do + run "abl_trainedA_nd${nd}_s${s}" --mode sdil --learn_A 1 --pert_ndirs $nd --depth 3 --width 256 --epochs $EP --seed $s + done + # E4 depth scaling + for d in 1 2 3 5 7; do + for m in bp dfa sdil; do + run "depth_${m}_d${d}_s${s}" --mode $m --depth $d --width 256 --epochs $EP --seed $s + done + done + # E5 feedback type + online apical control + for fb in error error_deriv; do + for ss in 0 5; do + run "ctrl_fb${fb}_settle${ss}_s${s}" --mode sdil --feedback $fb --settle_steps $ss \ + --kappa 0.3 --depth 3 --width 256 --epochs $EP --seed $s + done + done + # E6 predictor timescale (rho=2 so P matters); neutral vs task-period update + for neu in 0 1; do + for ep in 0.0005 0.002 0.01; do + run "pred_neu${neu}_etaP${ep}_s${s}" --mode sdil --nuis_rho 2 --learn_P 1 \ + --p_neutral $neu --eta_P $ep --depth 3 --width 256 --epochs $EP --seed $s + done + done + done +fi + +echo "=== battery done wave=$WAVE $(date) ===" >> "$PROG" diff --git a/experiments/run_v2.sh b/experiments/run_v2.sh new file mode 100644 index 0000000..1027ed7 --- /dev/null +++ b/experiments/run_v2.sh @@ -0,0 +1,19 @@ +#!/usr/bin/env bash +# Corrected sweep (sigma-scale fixed). Mode-specific LR. +# Usage: bash run_v2.sh <dataset> "<depths>" <residual> <act> "<seeds>" <ep> <pfx> +set -u; cd "$(dirname "$0")/.." +PY=/home/yurenh2/miniconda3/envs/ep_pascal/bin/python3 +export OMP_NUM_THREADS=2 +DS="$1"; DEPTHS="$2"; RES="$3"; ACT="$4"; SEEDS="$5"; EP="$6"; PFX="$7" +OUT=results; LOGS=logs/v2; mkdir -p "$OUT" "$LOGS"; PROG="$LOGS/progress.txt" +echo "=== v2 $(date) ds=$DS depths=[$DEPTHS] res=$RES act=$ACT seeds=[$SEEDS] ===" >> "$PROG" +run(){ local tag="$1"; shift; [ -f "$OUT/$tag.json" ] && { echo "skip $tag">>"$PROG"; return; } + echo ">>> $tag $(date +%H:%M:%S)">>"$PROG" + $PY experiments/run.py --tag "$tag" --outdir "$OUT" "$@" > "$LOGS/$tag.log" 2>&1 \ + && grep -h DONE "$LOGS/$tag.log">>"$PROG" || echo "FAIL $tag">>"$PROG"; } +for s in $SEEDS; do for d in $DEPTHS; do + run "${PFX}_${DS}_bp_d${d}_s${s}" --mode bp --dataset $DS --depth $d --width 256 --act $ACT --residual $RES --epochs $EP --seed $s --eta 0.05 + run "${PFX}_${DS}_dfa_d${d}_s${s}" --mode dfa --dataset $DS --depth $d --width 256 --act $ACT --residual $RES --epochs $EP --seed $s --eta 0.05 + run "${PFX}_${DS}_sdil_d${d}_s${s}" --mode sdil --dataset $DS --depth $d --width 256 --act $ACT --residual $RES --epochs $EP --seed $s --eta 0.02 --eta_A 0.02 --pert_ndirs 8 +done; done +echo "=== v2 done $(date) ===" >> "$PROG" diff --git a/experiments/seeds5.sh b/experiments/seeds5.sh new file mode 100644 index 0000000..f8ba51b --- /dev/null +++ b/experiments/seeds5.sh @@ -0,0 +1,4 @@ +cd /home/yurenh2/sdil +bash experiments/run_v2.sh fmnist "3 5 7" 0 tanh "1 2" 20 depth +bash experiments/run_v2.sh cifar10 "5 10" 1 tanh "1 2" 25 deepres +bash experiments/run_v2.sh mnist "3 5 7" 0 tanh "1 2" 20 depth diff --git a/experiments/sigfix.sh b/experiments/sigfix.sh new file mode 100644 index 0000000..9931e1d --- /dev/null +++ b/experiments/sigfix.sh @@ -0,0 +1,8 @@ +#!/usr/bin/env bash +cd /home/yurenh2/sdil; mkdir -p /tmp/sig logs/sig; SUM=logs/sig/summary.txt; : > "$SUM" +cf(){ local tag="$1"; shift; python3 experiments/run.py --device cpu --dataset fmnist --depth 5 --width 128 --epochs 5 --tag "$tag" --outdir /tmp/sig "$@" > logs/sig/$tag.log 2>&1; grep -h DONE logs/sig/$tag.log >> "$SUM"; } +cf s_dfa --mode dfa --eta 0.05 +cf s_sdil_e05 --mode sdil --eta 0.05 --pert_ndirs 8 +cf s_sdil_e02 --mode sdil --eta 0.02 --pert_ndirs 8 +cf s_sdil_e10 --mode sdil --eta 0.10 --pert_ndirs 8 +echo DONE_ALL >> "$SUM" diff --git a/experiments/smoke.py b/experiments/smoke.py new file mode 100644 index 0000000..f2efa87 --- /dev/null +++ b/experiments/smoke.py @@ -0,0 +1,86 @@ +"""Fast CPU correctness checks for SDIL mechanics and signs. +Run: python experiments/smoke.py +""" +import os +import sys +import torch + +sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) +from sdil.core import SDILNet, SDILConfig, sdil_step, node_perturbation_targets +from sdil.baselines import dfa_config +from sdil import probes +from sdil.data import get_dataset, onehot + + +def row_cos(u, v, eps=1e-12): + return ((u * v).sum(1) / (u.norm(dim=1) * v.norm(dim=1) + eps)).mean().item() + + +def main(): + torch.manual_seed(0) + dev = "cpu" + print("loading MNIST subset...") + tr, te, n_in, n_out = get_dataset("mnist", batch_size=128, device=dev) + xb, yb = next(iter(tr)) + x = xb[:256] + y = yb[:256] + yoh = onehot(y, 10) + + sizes = [784, 64, 64, 64, 10] + + # ---- CHECK 1: node-perturbation q estimates the descent direction -g ---- + net = SDILNet(sizes, device=dev, seed=1) + grads, _ = probes.true_hidden_grads(net, x, y) + for ndirs in (1, 8, 32): + qs = node_perturbation_targets(net, x, y, sigma=1e-2, n_dirs=ndirs) + cs = [row_cos(qs[l], -grads[l]) for l in range(len(qs))] + print(f"CHECK1 node-pert n_dirs={ndirs:2d} cos(q,-g) per layer = " + + " ".join(f"{c:+.3f}" for c in cs)) + assert all(c > 0 for c in cs), "node perturbation q should align with -grad" + + # ---- CHECK 2: SDIL overfits a fixed batch and alignment climbs ---- + net = SDILNet(sizes, device=dev, seed=2) + cfg = SDILConfig(eta=0.1, eta_A=0.05, eta_P=0.005, pert_every=2, pert_ndirs=8, + momentum=0.9) + print("\nCHECK2 SDIL overfit fixed batch:") + for step in range(401): + loss, _ = sdil_step(net, x, y, yoh, cfg, step) + if step % 50 == 0: + al = probes.alignment_report(net, x, y, yoh, cfg) + mc = sum(al["cos_r_negg"]) / len(al["cos_r_negg"]) + mca = sum(al["cos_apical_negg"]) / len(al["cos_apical_negg"]) + print(f" step {step:3d} loss {loss:.4f} mean_cos(r,-g) {mc:+.3f} " + f"mean_cos(apical,-g) {mca:+.3f} per-layer_r {['%.2f'%v for v in al['cos_r_negg']]}") + assert loss < 1.5, f"SDIL should reduce loss on fixed batch, got {loss}" + + # ---- CHECK 3: trained-A SDIL beats fixed-random-A (DFA) on alignment ---- + print("\nCHECK3 alignment: SDIL(trained A) vs DFA(fixed A):") + net_dfa = SDILNet(sizes, device=dev, seed=3) + cfg_dfa = dfa_config(eta=0.1, momentum=0.9) + for step in range(401): + sdil_step(net_dfa, x, y, yoh, cfg_dfa, step) + al_dfa = probes.alignment_report(net_dfa, x, y, yoh, cfg_dfa) + al_sdil = probes.alignment_report(net, x, y, yoh, cfg) + print(f" DFA mean_cos(r,-g) = {sum(al_dfa['cos_r_negg'])/len(al_dfa['cos_r_negg']):+.3f}") + print(f" SDIL mean_cos(r,-g) = {sum(al_sdil['cos_r_negg'])/len(al_sdil['cos_r_negg']):+.3f}") + + # ---- CHECK 4: nuisance -> residual preserves alignment, raw apical degrades ---- + print("\nCHECK4 residualization under soma-predictable nuisance (rho=2.0):") + net_n = SDILNet(sizes, device=dev, seed=4, nuis_rho=2.0) + cfg_n = SDILConfig(eta=0.1, eta_A=0.05, eta_P=0.02, pert_every=2, pert_ndirs=8, momentum=0.9) + for step in range(401): + sdil_step(net_n, x, y, yoh, cfg_n, step) + al_n = probes.alignment_report(net_n, x, y, yoh, cfg_n) + print(f" residual cos(r,-g) = {sum(al_n['cos_r_negg'])/len(al_n['cos_r_negg']):+.3f}") + print(f" raw apical cos(a,-g) = {sum(al_n['cos_apical_negg'])/len(al_n['cos_apical_negg']):+.3f}") + print(f" pure A c cos(Ac,-g) = {sum(al_n['cos_Ac_negg'])/len(al_n['cos_Ac_negg']):+.3f}") + + # ---- CHECK 5: single-step loss-decrease ratio vs exact GD ---- + print("\nCHECK5 single-step loss-decrease ratio vs exact GD:") + ldr = probes.loss_decrease_ratio(net, x, y, yoh, cfg, step=100) + print(f" dL_sdil={ldr['dL_sdil']:+.4f} dL_bp={ldr['dL_bp']:+.4f} ratio={ldr['ratio']:+.3f}") + print("\nALL SMOKE CHECKS PASSED") + + +if __name__ == "__main__": + main() diff --git a/experiments/zoo.py b/experiments/zoo.py new file mode 100644 index 0000000..8462be9 --- /dev/null +++ b/experiments/zoo.py @@ -0,0 +1,83 @@ +"""Train the local-learning baseline 'zoo' and report test accuracy, for the +comparison table alongside SDIL/DFA/BP. FA and EP are the working ones; PEPITA +and FF are included but currently undertuned. +Usage: python experiments/zoo.py --dataset mnist --methods fa,ep --epochs 10 +""" +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.local_baselines import FANet, PEPITANet, FFNet, EPNet +from sdil.baselines import evaluate +from sdil.data import get_dataset, onehot + + +def run_method(m, args, device): + tr, te, n_in, n_out = get_dataset(args.dataset, args.batch_size, device=device) + sizes = [n_in] + [args.width] * args.depth + [10] + t0 = time.time() + if m == "fa": + net = FANet(sizes, act="tanh", device=device, seed=args.seed) + for ep in range(args.epochs): + for x, y in tr: + net.fa_step(x, y, onehot(y, 10, device=device), args.eta, 0.9) + acc = evaluate(net, te)[0] + elif m == "pepita": + net = PEPITANet(sizes, act="tanh", device=device, seed=args.seed, f_scale=0.5) + for ep in range(args.epochs): + for x, y in tr: + net.pepita_step(x, y, onehot(y, 10, device=device), args.eta, 0.9) + acc = evaluate(net, te)[0] + elif m == "ff": + net = FFNet(sizes, act="relu", device=device, seed=args.seed, threshold=2.0, overlay_val=10.0) + for ep in range(args.epochs): + for x, y in tr: + net.train_step(x, y, args.eta) + acc = net.evaluate(te)[0] + elif m == "ep": + net = EPNet([n_in, args.width, 10], device=device, seed=args.seed, + beta=0.5, dt=0.5, T_free=20, T_nudge=8) + for ep in range(args.epochs): + for x, y in tr: + net.train_step(x, y, onehot(y, 10, device=device), args.eta) + acc = net.evaluate(te)[0] + else: + raise ValueError(m) + return acc, time.time() - t0 + + +def main(): + p = argparse.ArgumentParser() + p.add_argument("--dataset", default="mnist") + p.add_argument("--methods", default="fa,ep,pepita,ff") + p.add_argument("--depth", type=int, default=2) + p.add_argument("--width", type=int, default=500) + p.add_argument("--epochs", type=int, default=10) + p.add_argument("--batch_size", type=int, default=64) + p.add_argument("--eta", type=float, default=0.1) + 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") + p.add_argument("--tag", default="zoo") + args = p.parse_args() + + # sensible per-method default LRs if the shared one is off + lrs = {"fa": 0.05, "ep": 0.1, "pepita": 0.05, "ff": 0.03} + out = {"args": vars(args), "acc": {}} + for m in args.methods.split(","): + args.eta = lrs.get(m, args.eta) + acc, secs = run_method(m, args, args.device) + out["acc"][m] = acc + print(f"[{args.dataset}] {m}: test_acc {acc:.4f} ({secs:.0f}s)", flush=True) + os.makedirs(args.outdir, exist_ok=True) + with open(os.path.join(args.outdir, f"{args.tag}.json"), "w") as f: + json.dump(out, f) + print(f"saved -> {args.outdir}/{args.tag}.json", flush=True) + + +if __name__ == "__main__": + main() diff --git a/experiments/zoo_cpu.sh b/experiments/zoo_cpu.sh new file mode 100644 index 0000000..d98cac5 --- /dev/null +++ b/experiments/zoo_cpu.sh @@ -0,0 +1,6 @@ +cd /home/yurenh2/sdil +for ds in mnist fmnist; do + python3 experiments/zoo.py --dataset $ds --methods fa --depth 2 --width 256 --epochs 12 --tag zoo_fa_$ds >> logs/zoo.out 2>&1 + python3 experiments/zoo.py --dataset $ds --methods ep --depth 1 --width 500 --epochs 8 --tag zoo_ep_$ds >> logs/zoo.out 2>&1 +done +echo "ZOO DONE" >> logs/zoo.out |
