1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
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()
|