Port-Hamiltonian Neural ODE / stage2_bench.py

Mechanism confirmed, baseline not beaten

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