import json, random, sys import numpy as np import torch from torch import 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 = [1e-3, 3e-3, 1e-2] EPOCHS = 12 BATCH = 128 DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu' def seed_all(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def data(seed): ds = get_dataset('dynamics', seed=seed, n_train=400, n_test=200) return {k: (torch.as_tensor(v, dtype=torch.float32) if k in ('xtr','ytr','xte','yte') else v) for k, v in ds.items()} def baseline_run(lr, seed): seed_all(seed) ds = data(seed) model = make_model('rnn_small', tuple(ds['input_shape']), int(ds['out_dim'])) try: _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=lr, batch=BATCH, log=lambda *a, **k: None) except Exception: model = model.cpu() ds = {k: (v.cpu() if torch.is_tensor(v) else v) for k, v in ds.items()} _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=lr, batch=BATCH, log=lambda *a, **k: None) return float(metric) def hvp(loss, params, vector, create_graph=True): grads = torch.autograd.grad(loss, params, create_graph=True, allow_unused=True) dot = sum((g * v).sum() for g, v in zip(grads, vector) if g is not None) hv = torch.autograd.grad(dot, params, create_graph=create_graph, allow_unused=True) return [torch.zeros_like(p) if h is None else h for p, h in zip(params, hv)] def hess_regularizer(loss, model, nvec=2, eps=1e-4, kappa=10.0): params = flat_params = [p for p in model.parameters() if p.requires_grad] vals = [] for _ in range(nvec): vec = [torch.randn_like(p) for p in params] norm = torch.sqrt(sum((v * v).sum() for v in vec) + 1e-12) vec = [v / norm for v in vec] hv = hvp(loss, params, vec, create_graph=True) vals.append(sum((v * h).sum() for v, h in zip(vec, hv))) q = torch.stack(vals) penalty = torch.relu(-q + eps).square().mean() + 0.01 * torch.relu(q - kappa).square().mean() return penalty, q.detach() def idea_run(lr, seed, return_signature=False): seed_all(seed) ds = data(seed) dev = torch.device(DEVICE) model = make_model('rnn_small', tuple(ds['input_shape']), int(ds['out_dim'])).to(dev) x = ds['xtr'].to(dev); y = ds['ytr'].to(dev) xt = ds['xte'].to(dev); yt = ds['yte'].to(dev) opt = torch.optim.Adam(model.parameters(), lr=lr) loss_fn = nn.MSELoss() last_q = [] try: for _ in range(EPOCHS): order = torch.randperm(len(x), device=dev) for ix in order.split(BATCH): opt.zero_grad(set_to_none=True) pred = model(x[ix]) fit = loss_fn(pred, y[ix]) reg, q = hess_regularizer(fit, model) (fit + 0.05 * reg).backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0) opt.step() last_q = q.detach().cpu().tolist() with torch.no_grad(): metric = float(loss_fn(model(xt), yt).cpu()) except Exception: if dev.type == 'cuda': return idea_run_cpu(lr, seed, return_signature) raise if return_signature: return metric, model, ds, last_q return metric def idea_run_cpu(lr, seed, return_signature=False): global DEVICE old = DEVICE; DEVICE = 'cpu' try: return idea_run(lr, seed, return_signature) finally: DEVICE = old def flat_gradient(model, loss): return torch.cat([g.detach().flatten() for g in torch.autograd.grad(loss, model.parameters(), allow_unused=True) if g is not None]) def mechanism_signature(): seed = 9001 seed_all(seed) ds = data(seed) dev = torch.device(DEVICE) model = make_model('rnn_small', tuple(ds['input_shape']), int(ds['out_dim'])).to(dev) x, y = ds['xtr'][:64].to(dev), ds['ytr'][:64].to(dev) params = [p for p in model.parameters() if p.requires_grad] pred = model(x); loss = nn.MSELoss()(pred, y) g = torch.autograd.grad(loss, params) gnorm = float(torch.sqrt(sum((z*z).sum() for z in g)).detach().cpu()) # Rebuild the graph before the second-order response calculation. loss_h = nn.MSELoss()(model(x), y) vec = [torch.randn_like(p) for p in params] norm = torch.sqrt(sum((v*v).sum() for v in vec)); vec = [v/norm for v in vec] hv = hvp(loss_h, params, vec, create_graph=False) q = float(sum((v*h).sum() for v,h in zip(vec,hv)).detach().cpu()) observed = float(torch.sqrt(sum((h*h).sum() for h in hv)).detach().cpu()) predicted = abs(q) ratio = observed / (predicted + 1e-8) 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)} def main(): grid = [{'lr': lr} for lr in LRS] base = sweep_baseline(lambda cfg: (lambda seed: baseline_run(float(cfg['lr']), seed)), grid, seeds=SWEEP_SEEDS) trials = [{'cfg': c, 'result': evaluate(lambda seed, c=c: idea_run(float(c['lr']), seed), SEEDS)} for c in grid] best = min(trials, key=lambda z: z['result']['mean']) rep = make_report('dynamics', 'rnn_small', base, best['result'], {'idea_config':best['cfg'], 'idea_sweep':trials, 'mechanism_signature':mechanism_signature()}) rep['mechanism_signature'] = rep.pop('mechanism_signature') with open('bench_report.json','w') as f: json.dump(rep,f,indent=2) print(json.dumps(rep, indent=2)) if __name__ == '__main__': main()