Integrated-Growth Hopf Delay Scheduler / stage2_bench.py

Failed on benchmark

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