summaryrefslogtreecommitdiff
path: root/experiments/probe_credit.py
diff options
context:
space:
mode:
authorYurenHao0426 <Blackhao0426@gmail.com>2026-07-21 06:38:55 -0500
committerYurenHao0426 <Blackhao0426@gmail.com>2026-07-21 06:38:55 -0500
commit215acff518eea2d6d8a5bd7c90c398f49cc3b41b (patch)
tree8c79739d8b0a614ce76e685a3eff323d4cdb1742 /experiments/probe_credit.py
chore: capture initial SDIL project state
Diffstat (limited to 'experiments/probe_credit.py')
-rw-r--r--experiments/probe_credit.py114
1 files changed, 114 insertions, 0 deletions
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()