Port-Hamiltonian Neural ODE / stage2_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, time
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6import torch.nn.functional as F
7
8sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
9from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report, count_params
10
11SEED=2904
12EPOCHS=2
13BATCH=128
14# Union is used identically by both sides; baseline also sweeps method knob weight_decay.
15GRID=[{'lr':1e-3,'weight_decay':0.0},{'lr':3e-3,'weight_decay':0.0},{'lr':5e-3,'weight_decay':0.0}]
16
17class PHGRU(nn.Module):
18 """Port-Hamiltonian recurrent cell; sequence interface matches rnn_small."""
19 def __init__(self, hidden=4, eps=.03):
20 super().__init__(); self.hidden=hidden; self.eps=eps
21 self.inp=nn.Linear(3,hidden)
22 self.energy_net=nn.Sequential(nn.Linear(hidden,64),nn.Tanh(),nn.Linear(64,1))
23 self.a_net=nn.Sequential(nn.Linear(hidden,64),nn.Tanh(),nn.Linear(64,hidden*hidden))
24 self.l_net=nn.Sequential(nn.Linear(hidden,64),nn.Tanh(),nn.Linear(64,hidden*hidden))
25 self.readout=nn.Linear(hidden,1)
26 self.qraw=nn.Parameter(torch.eye(hidden)*.15)
27 def energy(self,z):
28 q=self.qraw@self.qraw.T + .01*torch.eye(self.hidden,device=z.device)
29 return F.softplus(self.energy_net(z).squeeze(-1)) + .5*((z@q)*z).sum(-1)
30 def vector_field(self,z):
31 # create_graph is needed during training because the learned energy gradient is differentiated
32 with torch.enable_grad():
33 zz=z.detach().requires_grad_(True); h=self.energy(zz)
34 g=torch.autograd.grad(h.sum(),zz,create_graph=self.training)[0]
35 b=z.shape[0]; A=self.a_net(zz).view(b,self.hidden,self.hidden)
36 L=self.l_net(zz).view(b,self.hidden,self.hidden)
37 J=A-A.transpose(1,2); R=L@L.transpose(1,2)+self.eps*torch.eye(self.hidden,device=z.device)
38 dz=torch.bmm((J-R),g.unsqueeze(-1)).squeeze(-1)
39 return dz, g, J, R
40 def forward(self,x, signature=False):
41 seq=x.view(x.shape[0],-1,3); z=torch.tanh(self.inp(seq[:,0]));
42 for k in range(1,seq.shape[1]):
43 u=seq[:,k]
44 dz,_,_,_=self.vector_field(z)
45 # input port: learned fixed B represented by input projection, preserving causal dynamics
46 z=z+0.12*dz+torch.tanh(self.inp(u))*0.05
47 z=torch.tanh(z)
48 out=self.readout(z)
49 if signature: return out,z
50 return out
51
52def seed_all(s):
53 np.random.seed(s); torch.manual_seed(s)
54
55def train_one(kind,cfg,seed,return_model=False):
56 seed_all(seed); ds=get_dataset('dynamics',seed,n_train=100,n_test=60)
57 model=make_model('rnn_small',ds['input_shape'],ds['out_dim']) if kind=='base' else PHGRU(hidden=4)
58 model,metric,hist=train_model(model,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg['weight_decay'],log=lambda *_:None)
59 if model is None: return (float('nan'), None, ds) if return_model else float('nan')
60 if return_model: return metric,model,ds
61 return metric
62
63def main():
64 # Cheap exact structural and energy identity check before training.
65 seed_all(SEED); m=PHGRU(hidden=6).eval(); z=torch.randn(32,6,requires_grad=True)
66 dz,g,J,R=m.vector_field(z); lhs=(g*dz).sum(1); rhs=-(g.unsqueeze(1)@[email protected](-1)).squeeze();
67 mathcheck={'max_skew':float((J+J.transpose(1,2)).abs().max()),'min_R_eig':float(torch.linalg.eigvalsh(R).amin()),'max_identity_residual':float((lhs-rhs).abs().max()),'max_dHdt':float(lhs.max())}
68 # Both systems see the exact same grid and paired seeds. Baseline sweep uses 4 seeds then full 8.
69 base=sweep_baseline(lambda cfg: (lambda s: train_one('base',cfg,s)),GRID)
70 full_grid=GRID
71 idea_cfg=min(GRID,key=lambda c: base['sweep'][GRID.index(c)]['mean'])
72 # evaluate idea at all three union configs; report best, while baseline was evaluated at every config.
73 idea_runs=[]
74 for cfg in full_grid:
75 r=evaluate(lambda s,cfg=cfg: train_one('idea',cfg,s))
76 idea_runs.append({'cfg':cfg,'result':r})
77 idea_best=min(idea_runs,key=lambda q:q['result']['mean'])
78 # Signature on trained models, measured behavior: observed dissipation identity and drift on NN states.
79 sig=[]
80 for s in range(8):
81 got=train_one('idea',idea_best['cfg'],s,True)
82 metric,model,ds=got
83 if model is None: continue
84 model.eval(); x=ds['xte'][:64].to(next(model.parameters()).device)
85 with torch.enable_grad():
86 z=torch.tanh(model.inp(x.view(x.shape[0],-1,3)[:,0])); dz,g,J,R=model.vector_field(z)
87 obs=(g*dz).sum(1); theo=-(g.unsqueeze(1)@[email protected](-1)).squeeze()
88 sig.append((float((J+J.transpose(1,2)).abs().max()),float(torch.linalg.eigvalsh(R).amin()),float((obs-theo).abs().max()),float(obs.max())))
89 signature={'trained_model_samples':len(sig),'max_skew_residual':max(x[0] for x in sig),'min_R_eigenvalue':min(x[1] for x in sig),'max_energy_identity_residual':max(x[2] for x in sig),'max_observed_dHdt':max(x[3] for x in sig),'predicted':{'skew_zero':True,'R_eigenvalue_ge_epsilon':True,'dHdt_le_zero':True},'confirmed':max(x[0] for x in sig)<1e-5 and min(x[1] for x in sig)>=.029 and max(x[2] for x in sig)<1e-4 and max(x[3] for x in sig)<=1e-5}
90 rep=make_report('dynamics','rnn_small',base,idea_best['result'],{'track_match':'stability/control -> dynamics','math_sanity':mathcheck,'mechanism_signature':signature,'idea_sweep':idea_runs,'selected_cfg':idea_best['cfg'],'parameter_counts':{'baseline':count_params(make_model('rnn_small',(8,3),1)),'idea':count_params(PHGRU(hidden=4))}})
91 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
92 print(json.dumps(rep,indent=2))
93if __name__=='__main__': main()