diff options
Diffstat (limited to 'experiments/probe_credit.py')
| -rw-r--r-- | experiments/probe_credit.py | 114 |
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() |
