summaryrefslogtreecommitdiff
path: root/experiments/probe_credit.py
blob: f7c553f8321c6e5fb59c95c8f6c844ffb4844840 (plain)
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()