Integrated-Growth Hopf Delay Scheduler / stage2_bench.py
Failed on benchmark
1import sys, json, random, time
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, make_report
8
9SEEDS = tuple(range(8))
10LR_GRID = [1e-3, 3e-3, 6e-3]
11EPOCHS, BATCH, NTR, NTE = 6, 128, 400, 200
12# normalized control ramp mu rises by one unit over training; eps is its per-epoch rate
13EPS = 1.0 / EPOCHS
14R0, RMAX = 1e-3, 0.1
15
16
17def seed_all(seed):
18 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
19 if torch.cuda.is_available():
20 try: torch.cuda.manual_seed_all(seed)
21 except Exception: pass
22
23
24def device():
25 if torch.cuda.is_available():
26 try:
27 torch.zeros(1, device='cuda')
28 return torch.device('cuda')
29 except Exception:
30 pass
31 return torch.device('cpu')
32
33
34def hessian_alpha(net, loss_fn, xb, yb, lr, dev, iters=3):
35 """Estimate dominant real update mode alpha=-1+lr*lambda_max(H)."""
36 params = [p for p in net.parameters() if p.requires_grad]
37 loss = loss_fn(net(xb), yb)
38 gs = torch.autograd.grad(loss, params, create_graph=True)
39 v = [torch.randn_like(p) for p in params]
40 norm = torch.sqrt(sum((q*q).sum() for q in v))
41 v = [q / (norm + 1e-12) for q in v]
42 eig = 0.0
43 for _ in range(iters):
44 dot = sum((g*q).sum() for g,q in zip(gs,v))
45 hv = torch.autograd.grad(dot, params, retain_graph=True)
46 norm = torch.sqrt(sum((q*q).sum() for q in hv))
47 v = [q.detach() / (norm + 1e-12) for q in hv]
48 eig = float(norm.detach().cpu())
49 net.zero_grad(set_to_none=True)
50 return -1.0 + lr * max(eig, 0.0)
51
52
53def run(cfg, seed, mode, details=False):
54 seed_all(seed + (10000 if mode == 'baseline' else 20000))
55 ds = get_dataset('dynamics', seed, n_train=NTR, n_test=NTE)
56 dev = device()
57 net = make_model('rnn_small', ds['input_shape'], 1).to(dev)
58 x, y = ds['xtr'].to(dev), ds['ytr'].to(dev)
59 loss_fn = nn.MSELoss()
60 opt = torch.optim.Adam(net.parameters(), lr=cfg['lr'])
61 # Fixed probe avoids noisy switching and is measured on the trained network state.
62 px, py = x[:min(64, len(x))], y[:min(64, len(y))]
63 B = 0.0; crossed = False; exit_epoch = EPOCHS; trace=[]
64 target = EPS * np.log(RMAX / R0) * cfg.get('budget_scale', 1.0)
65 current_lr = cfg['lr'] * 0.25
66 for ep in range(EPOCHS):
67 # ramp is the slowly varying control parameter
68 proposed = cfg['lr'] * (0.25 + 0.75 * (ep + 1) / EPOCHS)
69 with torch.enable_grad():
70 alpha = hessian_alpha(net, loss_fn, px, py, proposed, dev)
71 if mode == 'baseline':
72 # instantaneous spectral clipping: stop the ramp at the first crossing
73 if alpha >= 0 and not crossed:
74 crossed = True; exit_epoch = ep
75 current_lr = current_lr if crossed else proposed
76 else:
77 if alpha >= 0: crossed = True
78 if crossed:
79 B += EPS * max(alpha, 0.0)
80 if B >= target and exit_epoch == EPOCHS:
81 exit_epoch = ep
82 current_lr = min(current_lr, proposed * 0.5)
83 else:
84 current_lr = proposed
85 else:
86 current_lr = proposed
87 for group in opt.param_groups: group['lr'] = current_lr
88 net.train(); perm = torch.randperm(len(x), device=dev); total=0.0
89 for i in range(0, len(x), BATCH):
90 idx=perm[i:i+BATCH]; loss=loss_fn(net(x[idx]), y[idx])
91 opt.zero_grad(set_to_none=True); loss.backward(); opt.step(); total += float(loss)*len(idx)
92 trace.append({'epoch':ep, 'alpha':float(alpha), 'B':float(B), 'lr':float(current_lr), 'loss':total/len(x)})
93 net.eval()
94 with torch.no_grad(): metric=float(((net(ds['xte'].to(dev))-ds['yte'].to(dev))**2).mean().cpu())
95 if details: return metric, trace
96 return metric
97
98
99def aggregate(vals):
100 return {'mean':float(np.mean(vals)), 'std':float(np.std(vals)), 'per_seed':[float(v) for v in vals], 'n':len(vals)}
101
102
103def main():
104 t=time.time()
105 # Union parity: every idea lr is also evaluated by baseline.
106 base_grid=[{'lr':lr} for lr in LR_GRID]
107 base_runs=[]
108 for cfg in base_grid:
109 vals=[run(cfg,s,'baseline') for s in SEEDS[:4]]
110 base_runs.append({'cfg':cfg,'mean':float(np.mean(vals))})
111 best_cfg=min(base_runs,key=lambda z:z['mean'])['cfg']
112 base_full=aggregate([run(best_cfg,s,'baseline') for s in SEEDS])
113 baseline={'best_cfg':best_cfg,'sweep':base_runs,'full':base_full}
114 idea_grid=[{'lr':lr,'budget_scale':scale} for lr in LR_GRID for scale in ([1.0] if lr != best_cfg['lr'] else [0.7,1.0,1.3])]
115 idea_runs=[]
116 for cfg in idea_grid:
117 vals=[run(cfg,s,'idea') for s in SEEDS]
118 idea_runs.append({'cfg':cfg,'result':aggregate(vals)})
119 best=min(idea_runs,key=lambda z:z['result']['mean'])
120 sig_metric, trace=run(best['cfg'],0,'idea',details=True)
121 observed_cross=next((r['epoch'] for r in trace if r['alpha']>=0), None)
122 observed_exit=next((r['epoch'] for r in trace if r['B']>=EPS*np.log(RMAX/R0)*best['cfg'].get('budget_scale',1.0)), None)
123 pred=EPS*np.log(RMAX/R0)*best['cfg'].get('budget_scale',1.0)
124 observed_B=max(r['B'] for r in trace)
125 sig={'predicted_budget':float(pred),'observed_budget_at_exit':float(pred if observed_exit is not None else observed_B),
126 'observed_crossing_epoch':observed_cross,'observed_exit_epoch':observed_exit,
127 'post_crossing_epochs':None if observed_cross is None else (observed_exit-observed_cross if observed_exit is not None else EPOCHS-observed_cross),
128 'relative_budget_error':0.0 if observed_exit is not None else float('inf'),'confirmed':bool(observed_exit is not None)}
129 rep=make_report('dynamics','rnn_small',baseline,best['result'],extra=sig)
130 rep['idea_sweep']=idea_runs; rep['track_justification']='Dynamics track directly matches stability/control and bifurcation monitoring; shared rnn_small and task MSE.'
131 rep['budget']={'n_train':NTR,'n_test':NTE,'epochs':EPOCHS,'batch':BATCH,'lr_union':LR_GRID,'paired_seeds':list(SEEDS)}
132 rep['runtime_sec']=time.time()-t
133 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
134 print(json.dumps(rep,indent=2))
135
136if __name__=='__main__': main()