Removable-Pole Negative-Shifted Optimizer / stage2_negative_shift_bench.py
Failed on benchmark
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, sweep_baseline, evaluate, make_report
8
9SEEDS = tuple(range(8))
10LRS = [0.001, 0.003, 0.01]
11EPOCHS, BATCH = 18, 128
12NU_MULTS = [0.0, 0.05, 0.15]
13
14
15def seed_all(seed):
16 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
17 if torch.cuda.is_available():
18 try: torch.cuda.manual_seed_all(seed)
19 except Exception: pass
20
21
22def get_data(seed):
23 return get_dataset('tabular', seed, n_train=400, n_test=400)
24
25
26def hessian_scale(model, ds, dev):
27 """Power iteration on the minibatch Hessian, giving an actual curvature scale."""
28 model.zero_grad(set_to_none=True)
29 x, y = ds['xtr'][:128].to(dev), ds['ytr'][:128].to(dev)
30 loss = nn.MSELoss()(model(x), y)
31 gs = torch.autograd.grad(loss, tuple(model.parameters()), create_graph=True)
32 vs = [torch.randn_like(p) for p in model.parameters()]
33 norm = torch.sqrt(sum((v*v).sum() for v in vs))
34 vs = [v / norm for v in vs]
35 val = 1e-4
36 for _ in range(4):
37 dot = sum((g*v).sum() for g, v in zip(gs, vs))
38 hv = torch.autograd.grad(dot, tuple(model.parameters()), retain_graph=True)
39 norm = torch.sqrt(sum((h*h).sum() for h in hv))
40 val = float(norm.detach().cpu())
41 vs = [h / (norm + 1e-12) for h in hv]
42 model.zero_grad(set_to_none=True)
43 return max(val, 1e-5)
44
45
46def baseline_one(cfg, seed):
47 seed_all(seed); ds = get_data(seed)
48 model = make_model('mlp_tiny', ds['input_shape'], ds['out_dim'])
49 _, metric, _ = train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=BATCH,
50 weight_decay=cfg['weight_decay'], log=lambda *a: None)
51 return float('inf') if metric is None else metric
52
53
54def shifted_one(cfg, seed, collect=False):
55 seed_all(seed); ds = get_data(seed); model = make_model('mlp_tiny', ds['input_shape'], ds['out_dim'])
56 try:
57 dev = 'cuda' if torch.cuda.is_available() else 'cpu'
58 model = model.to(dev)
59 initial = [p.detach().clone() for p in model.parameters()]
60 scale = hessian_scale(model, ds, dev); nu = cfg['nu_mult'] * scale
61 opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'], weight_decay=cfg['weight_decay'])
62 x, y = ds['xtr'].to(dev), ds['ytr'].to(dev); lossf = nn.MSELoss(); hist = []
63 for _ in range(EPOCHS):
64 for ix in torch.randperm(len(x), device=dev).split(BATCH):
65 opt.zero_grad(set_to_none=True); loss = lossf(model(x[ix]), y[ix]); loss.backward()
66 # Apply the negative quadratic term to displacement delta=p-p0.
67 with torch.no_grad():
68 for p, p0 in zip(model.parameters(), initial):
69 if p.grad is not None: p.grad.add_(-nu * (p - p0))
70 torch.nn.utils.clip_grad_norm_(model.parameters(), 10.0); opt.step()
71 hist.append(float(loss.detach().cpu()))
72 with torch.no_grad():
73 test = float(lossf(model(ds['xte'].to(dev)), ds['yte'].to(dev)).cpu())
74 disp = float(torch.sqrt(sum(((p-p0)**2).sum() for p,p0 in zip(model.parameters(), initial))).cpu())
75 if collect: return test, {'model': model, 'ds': ds, 'scale': scale, 'nu': nu, 'disp': disp, 'history': hist}
76 return test
77 except Exception:
78 try: torch.cuda.empty_cache()
79 except Exception: pass
80 return float('inf')
81
82
83def main():
84 base_grid = [{'lr': lr, 'weight_decay': wd} for lr in LRS for wd in [0.0, 1e-4]]
85 baseline = sweep_baseline(lambda c: lambda s: baseline_one(c, s), base_grid, seeds=(0,1,2,3))
86 # Same lr union and baseline's selected central regularization; only nu differs.
87 wd = baseline['best_cfg']['weight_decay']
88 idea_grid = [{'lr': lr, 'weight_decay': wd, 'nu_mult': nu} for lr in LRS for nu in NU_MULTS]
89 tried = []
90 for cfg in idea_grid:
91 r = evaluate(lambda s, c=cfg: shifted_one(c, s), seeds=(0,1,2,3))
92 tried.append({'cfg': cfg, 'mean': r['mean']})
93 best_cfg = min(idea_grid, key=lambda c: next(q['mean'] for q in tried if q['cfg'] == c))
94 idea = evaluate(lambda s: shifted_one(best_cfg, s), seeds=SEEDS)
95 # Signature is measured on trained models: positive-shift run versus zero-shift run.
96 sig_cfg = next(c for c in idea_grid if c['lr'] == best_cfg['lr'] and c['nu_mult'] == 0.15)
97 a = shifted_one(sig_cfg, 0, collect=True); z = shifted_one({**sig_cfg, 'nu_mult': 0.0}, 0, collect=True)
98 predicted = 1.0 + sig_cfg['lr'] * a[1]['nu']
99 observed = (a[1]['disp'] / max(z[1]['disp'], 1e-12))
100 signature = {'predicted': {'one_step_displacement_factor': predicted, 'nu': a[1]['nu'], 'curvature_scale': a[1]['scale']},
101 'observed': {'trained_positive_shift_disp': a[1]['disp'], 'trained_zero_shift_disp': z[1]['disp'], 'whole_run_ratio': observed},
102 'confirmed': bool(np.isfinite(observed) and abs(observed-predicted) / max(abs(predicted),1e-9) < 0.25)}
103 rep = make_report('tabular', 'mlp_tiny', baseline, idea, signature)
104 rep['idea_sweep'] = {'grid': tried, 'best_cfg': best_cfg, 'shared_lr_union': LRS}
105 Path('bench_report.json').write_text(json.dumps(rep, indent=2)); print(json.dumps(rep, indent=2))
106
107if __name__ == '__main__': main()