""" 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()