Sensitivity-Conditioned Neural ODE Pruning / bench_sensitivity_pruning.py
Failed on benchmark
1import json, random, sys
2import numpy as np
3import torch
4import torch.nn as nn
5
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
8
9SEEDS = tuple(range(8))
10SWEEP_SEEDS = tuple(range(4))
11LRS = (0.0015, 0.003, 0.006)
12EPOCHS = 15
13NTR, NTE = 400, 200
14KEEP = 32
15
16
17def seed_all(seed):
18 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
19 if torch.cuda.is_available():
20 try: torch.cuda.manual_seed_all(seed)
21 except Exception: pass
22
23
24def masked_model(base, indices):
25 """A post-training pruned rnn_small: selected hidden units remain active.
26 The recurrent architecture and training loop are otherwise unchanged."""
27 class Masked(nn.Module):
28 def __init__(self, b, idx):
29 super().__init__(); self.base = b
30 m = torch.zeros(b.rnn.hidden_size)
31 m[list(idx)] = 1.0
32 self.register_buffer('mask', m)
33 def forward(self, x):
34 seq = x.view(x.shape[0], -1, 3)
35 try:
36 _, h = self.base.rnn(seq)
37 except RuntimeError:
38 cudnn = torch.backends.cudnn.enabled; torch.backends.cudnn.enabled = False
39 try: _, h = self.base.rnn(seq)
40 finally: torch.backends.cudnn.enabled = cudnn
41 return self.base.head(h[-1] * self.mask.view(1, -1))
42 return Masked(base, indices)
43
44
45def gru_hidden(net, seq):
46 try:
47 return net.rnn(seq)[1]
48 except RuntimeError:
49 old = torch.backends.cudnn.enabled
50 torch.backends.cudnn.enabled = False
51 try:
52 return net.rnn(seq)[1]
53 finally:
54 torch.backends.cudnn.enabled = old
55
56
57def scores(net, ds):
58 """J_g is the trained model's observed-output sensitivity to head group g."""
59 net.eval(); dev = next(net.parameters()).device
60 x = ds['xtr'].to(dev)
61 with torch.no_grad():
62 seq = x.view(x.shape[0], -1, 3)
63 h = gru_hidden(net, seq)
64 J = h[-1].detach().cpu().numpy()
65 J = J - J.mean(0, keepdims=True)
66 info = np.sum(J * J, axis=0)
67 residual = np.zeros(J.shape[1])
68 for g in range(J.shape[1]):
69 other = np.delete(J, g, axis=1)
70 q, _ = np.linalg.qr(other, mode='reduced')
71 z = J[:, g:g+1] - q @ (q.T @ J[:, g:g+1])
72 residual[g] = np.sum(z*z)
73 # prioritize both observable information and incremental rank
74 score = residual * np.sqrt(info + 1e-12) / (info + 1e-12)
75 return info, residual, score
76
77
78def choose(net, ds, kind):
79 if kind == 'sensitivity':
80 info, residual, score = scores(net, ds)
81 idx = np.argsort(score)[-KEEP:]
82 return np.sort(idx), info, residual
83 # standard magnitude pruning: outgoing head weights define hidden-unit magnitude
84 w = net.head.weight.detach().cpu().numpy()
85 mag = np.linalg.norm(w, axis=0)
86 idx = np.argsort(mag)[-KEEP:]
87 info, residual, _ = scores(net, ds)
88 return np.sort(idx), info, residual
89
90
91def run_one(seed, lr, kind):
92 seed_all(seed)
93 ds = get_dataset('dynamics', seed, n_train=NTR, n_test=NTE)
94 model = make_model('rnn_small', ds['input_shape'], ds['out_dim'])
95 full, _, _ = train_model(model, ds, epochs=EPOCHS, lr=lr, batch=128, log=lambda *_: None)
96 if full is None: return float('nan'), {}
97 idx, info, residual = choose(full, ds, kind)
98 pruned = masked_model(full, idx)
99 # Retraining is required by the proposed iterative pruning procedure.
100 net, metric, _ = train_model(pruned, ds, epochs=EPOCHS, lr=lr, batch=128, log=lambda *_: None)
101 if net is None: return float('nan'), {}
102 with torch.no_grad():
103 pred = net(ds['xte'].to(next(net.parameters()).device)).detach().cpu().numpy().ravel()
104 # trained-model behavior signature: score-predicted importance vs observed ablation
105 dev = next(net.parameters()).device
106 x = ds['xte'].to(dev)
107 with torch.no_grad():
108 seq = x.view(x.shape[0], -1, 3); h = gru_hidden(net.base, seq); h = h[-1]
109 y0 = net.base.head(h * net.mask.view(1,-1))
110 observed = []
111 for g in range(64):
112 mm = net.mask.clone(); mm[g] = 0
113 observed.append(float(torch.mean((y0-net.base.head(h*mm.view(1,-1)))**2).cpu()))
114 observed = np.asarray(observed)
115 corr = float(np.corrcoef(info, observed)[0,1]) if np.std(info)>0 and np.std(observed)>0 else 0.0
116 return float(metric), {'retained_units': int(len(idx)), 'params_total': int(sum(p.numel() for p in net.parameters())), 'mean_information_retained': float(np.mean(info[idx])), 'mean_residual_retained': float(np.mean(residual[idx])), 'importance_ablation_corr': corr, 'predicted_information': float(np.mean(info)), 'observed_ablation': float(np.mean(observed))}
117
118
119def make_train(kind, lr):
120 return lambda seed: run_one(seed, lr, kind)[0]
121
122
123def main():
124 # Every idea learning rate is also evaluated by baseline, satisfying union parity.
125 grid = [{'lr': lr, 'epochs': EPOCHS, 'keep': KEEP, 'method': 'magnitude'} for lr in LRS]
126 base = sweep_baseline(lambda cfg: make_train('magnitude', cfg['lr']), grid, seeds=SWEEP_SEEDS)
127 idea_trials = []
128 for lr in LRS:
129 r = evaluate(make_train('sensitivity', lr), seeds=SEEDS)
130 idea_trials.append({'cfg': {'lr': lr, 'epochs': EPOCHS, 'keep': KEEP, 'method': 'sensitivity'}, **r})
131 best = min(idea_trials, key=lambda x: x['mean'])
132 # Collect per-seed mechanism values for the selected configuration.
133 sig = [run_one(s, best['cfg']['lr'], 'sensitivity')[1] for s in SEEDS]
134 signature = {'prediction': 'weighted sensitivity information and incremental residual rank identify useful hidden groups', 'per_seed': sig, 'confirmed': bool(np.mean([x.get('importance_ablation_corr', 0) for x in sig]) > 0.5)}
135 idea = {'best_cfg': best['cfg'], 'sweep': [{'cfg': x['cfg'], 'mean': x['mean']} for x in idea_trials], 'mean': best['mean'], 'std': best['std'], 'per_seed': best['per_seed'], 'n': best['n']}
136 report = make_report('dynamics', 'rnn_small', base, idea, signature)
137 report['protocol_notes'] = {'paired_seeds': list(SEEDS), 'n_train': NTR, 'n_test': NTE, 'retained_units': KEEP, 'effective_unit_fraction': KEEP/64.0, 'baseline_method': 'outgoing-head magnitude selection', 'idea_method': 'trained-output sensitivity trace plus leave-one-group residual selection', 'custom_track': None}
138 with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2)
139 print(json.dumps(report, indent=2))
140
141if __name__ == '__main__': main()