import json, random, sys import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report SEEDS = tuple(range(8)) SWEEP_SEEDS = tuple(range(4)) LRS = (0.0015, 0.003, 0.006) EPOCHS = 15 NTR, NTE = 400, 200 KEEP = 32 def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): try: torch.cuda.manual_seed_all(seed) except Exception: pass def masked_model(base, indices): """A post-training pruned rnn_small: selected hidden units remain active. The recurrent architecture and training loop are otherwise unchanged.""" class Masked(nn.Module): def __init__(self, b, idx): super().__init__(); self.base = b m = torch.zeros(b.rnn.hidden_size) m[list(idx)] = 1.0 self.register_buffer('mask', m) def forward(self, x): seq = x.view(x.shape[0], -1, 3) try: _, h = self.base.rnn(seq) except RuntimeError: cudnn = torch.backends.cudnn.enabled; torch.backends.cudnn.enabled = False try: _, h = self.base.rnn(seq) finally: torch.backends.cudnn.enabled = cudnn return self.base.head(h[-1] * self.mask.view(1, -1)) return Masked(base, indices) def gru_hidden(net, seq): try: return net.rnn(seq)[1] except RuntimeError: old = torch.backends.cudnn.enabled torch.backends.cudnn.enabled = False try: return net.rnn(seq)[1] finally: torch.backends.cudnn.enabled = old def scores(net, ds): """J_g is the trained model's observed-output sensitivity to head group g.""" net.eval(); dev = next(net.parameters()).device x = ds['xtr'].to(dev) with torch.no_grad(): seq = x.view(x.shape[0], -1, 3) h = gru_hidden(net, seq) J = h[-1].detach().cpu().numpy() J = J - J.mean(0, keepdims=True) info = np.sum(J * J, axis=0) residual = np.zeros(J.shape[1]) for g in range(J.shape[1]): other = np.delete(J, g, axis=1) q, _ = np.linalg.qr(other, mode='reduced') z = J[:, g:g+1] - q @ (q.T @ J[:, g:g+1]) residual[g] = np.sum(z*z) # prioritize both observable information and incremental rank score = residual * np.sqrt(info + 1e-12) / (info + 1e-12) return info, residual, score def choose(net, ds, kind): if kind == 'sensitivity': info, residual, score = scores(net, ds) idx = np.argsort(score)[-KEEP:] return np.sort(idx), info, residual # standard magnitude pruning: outgoing head weights define hidden-unit magnitude w = net.head.weight.detach().cpu().numpy() mag = np.linalg.norm(w, axis=0) idx = np.argsort(mag)[-KEEP:] info, residual, _ = scores(net, ds) return np.sort(idx), info, residual def run_one(seed, lr, kind): seed_all(seed) ds = get_dataset('dynamics', seed, n_train=NTR, n_test=NTE) model = make_model('rnn_small', ds['input_shape'], ds['out_dim']) full, _, _ = train_model(model, ds, epochs=EPOCHS, lr=lr, batch=128, log=lambda *_: None) if full is None: return float('nan'), {} idx, info, residual = choose(full, ds, kind) pruned = masked_model(full, idx) # Retraining is required by the proposed iterative pruning procedure. net, metric, _ = train_model(pruned, ds, epochs=EPOCHS, lr=lr, batch=128, log=lambda *_: None) if net is None: return float('nan'), {} with torch.no_grad(): pred = net(ds['xte'].to(next(net.parameters()).device)).detach().cpu().numpy().ravel() # trained-model behavior signature: score-predicted importance vs observed ablation dev = next(net.parameters()).device x = ds['xte'].to(dev) with torch.no_grad(): seq = x.view(x.shape[0], -1, 3); h = gru_hidden(net.base, seq); h = h[-1] y0 = net.base.head(h * net.mask.view(1,-1)) observed = [] for g in range(64): mm = net.mask.clone(); mm[g] = 0 observed.append(float(torch.mean((y0-net.base.head(h*mm.view(1,-1)))**2).cpu())) observed = np.asarray(observed) corr = float(np.corrcoef(info, observed)[0,1]) if np.std(info)>0 and np.std(observed)>0 else 0.0 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))} def make_train(kind, lr): return lambda seed: run_one(seed, lr, kind)[0] def main(): # Every idea learning rate is also evaluated by baseline, satisfying union parity. grid = [{'lr': lr, 'epochs': EPOCHS, 'keep': KEEP, 'method': 'magnitude'} for lr in LRS] base = sweep_baseline(lambda cfg: make_train('magnitude', cfg['lr']), grid, seeds=SWEEP_SEEDS) idea_trials = [] for lr in LRS: r = evaluate(make_train('sensitivity', lr), seeds=SEEDS) idea_trials.append({'cfg': {'lr': lr, 'epochs': EPOCHS, 'keep': KEEP, 'method': 'sensitivity'}, **r}) best = min(idea_trials, key=lambda x: x['mean']) # Collect per-seed mechanism values for the selected configuration. sig = [run_one(s, best['cfg']['lr'], 'sensitivity')[1] for s in SEEDS] 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)} 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']} report = make_report('dynamics', 'rnn_small', base, idea, signature) 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} with open('bench_report.json', 'w') as f: json.dump(report, f, indent=2) print(json.dumps(report, indent=2)) if __name__ == '__main__': main()