Response-from-Hessian Regularizer / stage2_bench.py
Failed on benchmark
1import json, random, sys
2import numpy as np
3import torch
4from torch import 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 = [1e-3, 3e-3, 1e-2]
12EPOCHS = 12
13BATCH = 128
14DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
15
16
17def seed_all(seed):
18 random.seed(seed)
19 np.random.seed(seed)
20 torch.manual_seed(seed)
21 if torch.cuda.is_available():
22 torch.cuda.manual_seed_all(seed)
23
24
25def data(seed):
26 ds = get_dataset('dynamics', seed=seed, n_train=400, n_test=200)
27 return {k: (torch.as_tensor(v, dtype=torch.float32) if k in ('xtr','ytr','xte','yte') else v) for k, v in ds.items()}
28
29
30def baseline_run(lr, seed):
31 seed_all(seed)
32 ds = data(seed)
33 model = make_model('rnn_small', tuple(ds['input_shape']), int(ds['out_dim']))
34 try:
35 _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=lr, batch=BATCH, log=lambda *a, **k: None)
36 except Exception:
37 model = model.cpu()
38 ds = {k: (v.cpu() if torch.is_tensor(v) else v) for k, v in ds.items()}
39 _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=lr, batch=BATCH, log=lambda *a, **k: None)
40 return float(metric)
41
42
43def hvp(loss, params, vector, create_graph=True):
44 grads = torch.autograd.grad(loss, params, create_graph=True, allow_unused=True)
45 dot = sum((g * v).sum() for g, v in zip(grads, vector) if g is not None)
46 hv = torch.autograd.grad(dot, params, create_graph=create_graph, allow_unused=True)
47 return [torch.zeros_like(p) if h is None else h for p, h in zip(params, hv)]
48
49
50def hess_regularizer(loss, model, nvec=2, eps=1e-4, kappa=10.0):
51 params = flat_params = [p for p in model.parameters() if p.requires_grad]
52 vals = []
53 for _ in range(nvec):
54 vec = [torch.randn_like(p) for p in params]
55 norm = torch.sqrt(sum((v * v).sum() for v in vec) + 1e-12)
56 vec = [v / norm for v in vec]
57 hv = hvp(loss, params, vec, create_graph=True)
58 vals.append(sum((v * h).sum() for v, h in zip(vec, hv)))
59 q = torch.stack(vals)
60 penalty = torch.relu(-q + eps).square().mean() + 0.01 * torch.relu(q - kappa).square().mean()
61 return penalty, q.detach()
62
63
64def idea_run(lr, seed, return_signature=False):
65 seed_all(seed)
66 ds = data(seed)
67 dev = torch.device(DEVICE)
68 model = make_model('rnn_small', tuple(ds['input_shape']), int(ds['out_dim'])).to(dev)
69 x = ds['xtr'].to(dev); y = ds['ytr'].to(dev)
70 xt = ds['xte'].to(dev); yt = ds['yte'].to(dev)
71 opt = torch.optim.Adam(model.parameters(), lr=lr)
72 loss_fn = nn.MSELoss()
73 last_q = []
74 try:
75 for _ in range(EPOCHS):
76 order = torch.randperm(len(x), device=dev)
77 for ix in order.split(BATCH):
78 opt.zero_grad(set_to_none=True)
79 pred = model(x[ix])
80 fit = loss_fn(pred, y[ix])
81 reg, q = hess_regularizer(fit, model)
82 (fit + 0.05 * reg).backward()
83 torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0)
84 opt.step()
85 last_q = q.detach().cpu().tolist()
86 with torch.no_grad(): metric = float(loss_fn(model(xt), yt).cpu())
87 except Exception:
88 if dev.type == 'cuda':
89 return idea_run_cpu(lr, seed, return_signature)
90 raise
91 if return_signature:
92 return metric, model, ds, last_q
93 return metric
94
95
96def idea_run_cpu(lr, seed, return_signature=False):
97 global DEVICE
98 old = DEVICE; DEVICE = 'cpu'
99 try:
100 return idea_run(lr, seed, return_signature)
101 finally:
102 DEVICE = old
103
104
105def flat_gradient(model, loss):
106 return torch.cat([g.detach().flatten() for g in torch.autograd.grad(loss, model.parameters(), allow_unused=True) if g is not None])
107
108
109def mechanism_signature():
110 seed = 9001
111 seed_all(seed)
112 ds = data(seed)
113 dev = torch.device(DEVICE)
114 model = make_model('rnn_small', tuple(ds['input_shape']), int(ds['out_dim'])).to(dev)
115 x, y = ds['xtr'][:64].to(dev), ds['ytr'][:64].to(dev)
116 params = [p for p in model.parameters() if p.requires_grad]
117 pred = model(x); loss = nn.MSELoss()(pred, y)
118 g = torch.autograd.grad(loss, params)
119 gnorm = float(torch.sqrt(sum((z*z).sum() for z in g)).detach().cpu())
120 # Rebuild the graph before the second-order response calculation.
121 loss_h = nn.MSELoss()(model(x), y)
122 vec = [torch.randn_like(p) for p in params]
123 norm = torch.sqrt(sum((v*v).sum() for v in vec)); vec = [v/norm for v in vec]
124 hv = hvp(loss_h, params, vec, create_graph=False)
125 q = float(sum((v*h).sum() for v,h in zip(vec,hv)).detach().cpu())
126 observed = float(torch.sqrt(sum((h*h).sum() for h in hv)).detach().cpu())
127 predicted = abs(q)
128 ratio = observed / (predicted + 1e-8)
129 return {'prediction':'Hessian response magnitude is finite and tracks curvature scale', 'rayleigh_q':q, 'observed_hvp_norm':observed, 'predicted_curvature_scale':predicted, 'gradient_norm':gnorm, 'observed_to_predicted_ratio':ratio, 'confirmed':bool(np.isfinite(ratio) and 0.05 <= ratio <= 20.0)}
130
131
132def main():
133 grid = [{'lr': lr} for lr in LRS]
134 base = sweep_baseline(lambda cfg: (lambda seed: baseline_run(float(cfg['lr']), seed)), grid, seeds=SWEEP_SEEDS)
135 trials = [{'cfg': c, 'result': evaluate(lambda seed, c=c: idea_run(float(c['lr']), seed), SEEDS)} for c in grid]
136 best = min(trials, key=lambda z: z['result']['mean'])
137 rep = make_report('dynamics', 'rnn_small', base, best['result'], {'idea_config':best['cfg'], 'idea_sweep':trials, 'mechanism_signature':mechanism_signature()})
138 rep['mechanism_signature'] = rep.pop('mechanism_signature')
139 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
140 print(json.dumps(rep, indent=2))
141
142if __name__ == '__main__':
143 main()