summaryrefslogtreecommitdiff
path: root/experiments
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
chore: capture initial SDIL project state
Diffstat (limited to 'experiments')
-rw-r--r--experiments/analyze.py191
-rw-r--r--experiments/analyze_depth.py134
-rw-r--r--experiments/calib_tent.py18
-rw-r--r--experiments/deep_sweep.sh37
-rw-r--r--experiments/depth_sweep.sh39
-rw-r--r--experiments/diag_norm.sh15
-rw-r--r--experiments/probe_credit.py114
-rw-r--r--experiments/run.py177
-rw-r--r--experiments/run_battery.sh88
-rw-r--r--experiments/run_v2.sh19
-rw-r--r--experiments/seeds5.sh4
-rw-r--r--experiments/sigfix.sh8
-rw-r--r--experiments/smoke.py86
-rw-r--r--experiments/zoo.py83
-rw-r--r--experiments/zoo_cpu.sh6
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