Response-from-Hessian Regularizer / stage2_bench.py

Failed on benchmark

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