Sensitivity-Conditioned Neural ODE Pruning / bench_sensitivity_pruning.py

Failed on benchmark

Raw ⬇ ZIP
  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()