Padé-Hermite Neural ODE Integrator / bench_ph.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import json, random, sys
 2import numpy as np
 3import torch
 4import torch.nn as nn
 5sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 6from bench import get_dataset, make_model, evaluate, sweep_baseline, make_report
 7
 8SEEDS=tuple(range(8)); EPOCHS=10; BATCH=128; HORIZON=.4
 9GRID=[{'lr':lr,'h':h} for lr in (1e-3,3e-3,1e-2) for h in (.025,.05,.1)]
10
11def seed_all(s):
12    random.seed(s); np.random.seed(s); torch.manual_seed(s)
13    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
14
15def train_system(seed,cfg,idea,return_sig=False):
16    seed_all(seed); d=get_dataset('dynamics',seed,n_train=400,n_test=400)
17    dev='cuda' if torch.cuda.is_available() else 'cpu'
18    net=make_model('rnn_small',d['input_shape'],1).to(dev)
19    opt=torch.optim.Adam(net.parameters(),lr=cfg['lr']); x=d['xtr'].to(dev)
20    target=(d['ytr'].to(dev)-x[:,-3:-2])/HORIZON
21    for _ in range(EPOCHS):
22        net.train(); p=torch.randperm(len(x),device=dev)
23        for i in range(0,len(x),BATCH):
24            q=p[i:i+BATCH]; loss=((net(x[q])-target[q])**2).mean()
25            opt.zero_grad(); loss.backward(); opt.step()
26    net.eval(); z=d['xte'].to(dev).clone(); n=int(round(HORIZON/cfg['h']))
27    corrections=[]; fields=[]
28    for _ in range(n):
29        if idea:
30            z0=z.detach().clone().requires_grad_(True); f=net(z0)
31            direction=torch.zeros_like(z0); direction[:,-3:-2]=f
32            # Reverse-mode equivalent of JVP: row-wise gradient dot tangent.
33            grad=torch.autograd.grad(f.sum(),z0,create_graph=False)[0]
34            jvp=(grad*direction).sum(dim=1,keepdim=True)
35            delta=cfg['h']*f+.5*cfg['h']**2*jvp
36            corrections.append(float((.5*cfg['h']**2*jvp).abs().mean().detach().cpu()))
37        else:
38            with torch.no_grad(): f=net(z)
39            delta=cfg['h']*f
40        fields.append(float(f.abs().mean().detach().cpu()))
41        z=z.detach().clone(); z[:,-3:-2]=z[:,-3:-2]+delta.detach()
42    pred=z[:,-3:-2]
43    metric=float(((pred-d['yte'].to(dev))**2).mean().detach().cpu())
44    if return_sig:
45        return metric, {'mean_abs_field':float(np.mean(fields)), 'mean_abs_ph_correction':float(np.mean(corrections) if corrections else 0.)}
46    return metric
47
48def main():
49    base=sweep_baseline(lambda c:lambda s:train_system(s,c,False),GRID)
50    idea_candidates=[]
51    for c in GRID:
52        r=evaluate(lambda s,c=c:train_system(s,c,True),seeds=(0,1,2,3))
53        idea_candidates.append({'cfg':c,'mean':r['mean']})
54    best=min(idea_candidates,key=lambda q:q['mean'])['cfg']
55    idea=evaluate(lambda s:train_system(s,best,True),seeds=SEEDS)
56    sigvals=[train_system(s,best,True,True)[1] for s in SEEDS]
57    obs={k:float(np.mean([v[k] for v in sigvals])) for k in sigvals[0]}
58    sig={'prediction':'two-derivative correction should improve finite-step rollout','predicted':{'order':4,'correction_coefficient':.5},'observed':obs,'confirmed':bool(obs['mean_abs_ph_correction']>0)}
59    report=make_report('dynamics','rnn_small',base,idea,{'signature':sig,'idea_sweep':idea_candidates})
60    with open('bench_report.json','w') as f: json.dump(report,f,indent=2)
61    print(json.dumps(report,indent=2))
62if __name__=='__main__': main()