Port-Hamiltonian Neural ODE / experiment.py
Mechanism confirmed, baseline not beaten
1import json, math, time
2import numpy as np
3import torch
4import torch.nn as nn
5
6SEED = 2904
7torch.manual_seed(SEED); np.random.seed(SEED)
8torch.set_default_dtype(torch.float64)
9DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
10try:
11 if DEVICE == "cuda":
12 torch.cuda.empty_cache()
13except Exception:
14 DEVICE = "cpu"
15
16class MLP(nn.Module):
17 def __init__(self, inp, out, width=32):
18 super().__init__()
19 self.net = nn.Sequential(nn.Linear(inp,width), nn.Tanh(), nn.Linear(width,width), nn.Tanh(), nn.Linear(width,out))
20 def forward(self,x): return self.net(x)
21
22class PHField(nn.Module):
23 def __init__(self, d=2, eps=0.03):
24 super().__init__(); self.d=d; self.eps=eps
25 self.h=MLP(d,1); self.a=MLP(d,d*d); self.l=MLP(d,d*d)
26 self.qraw=nn.Parameter(torch.eye(d)*0.7)
27 def energy(self,z):
28 q=self.qraw @ self.qraw.T + 0.05*torch.eye(self.d, device=z.device)
29 return torch.nn.functional.softplus(self.h(z).squeeze(-1)) + .5*torch.sum((z@q)*z,dim=-1)
30 def matrices(self,z):
31 n=z.shape[0]; A=self.a(z).reshape(n,self.d,self.d); L=self.l(z).reshape(n,self.d,self.d)
32 J=A-A.transpose(1,2); R=L@L.transpose(1,2)+self.eps*torch.eye(self.d,device=z.device)
33 return J,R
34 def forward(self,z):
35 zz=z.detach().requires_grad_(True)
36 H=self.energy(zz); g=torch.autograd.grad(H.sum(),zz,create_graph=self.training)[0]
37 J,R=self.matrices(zz)
38 return torch.bmm((J-R),g.unsqueeze(-1)).squeeze(-1)
39 def derivative_terms(self,z):
40 zz=z.detach().requires_grad_(True); H=self.energy(zz); g=torch.autograd.grad(H.sum(),zz)[0]
41 J,R=self.matrices(zz); dz=torch.bmm((J-R),g.unsqueeze(-1)).squeeze(-1)
42 return H.detach(), g.detach(), J.detach(), R.detach(), (g*dz).sum(1).detach(), -(torch.bmm(R,g.unsqueeze(-1)).squeeze(-1)*g).sum(1).detach()
43
44class Unconstrained(nn.Module):
45 def __init__(self,d=2): super().__init__(); self.net=MLP(d,d)
46 def forward(self,z): return self.net(z)
47
48def rk4(fun,x,dt):
49 k1=fun(x); k2=fun(x+.5*dt*k1); k3=fun(x+.5*dt*k2); k4=fun(x+dt*k3)
50 return x+dt*(k1+2*k2+2*k3+k4)/6
51
52def oscillator(x):
53 # nonlinear Hamiltonian oscillator with linear damping
54 q,p=x[...,0],x[...,1]
55 return torch.stack((p, -q-0.15*q**3-0.12*p),-1)
56
57def make_data(n=32, steps=16, dt=.08):
58 x=torch.randn(n,2)*1.1; xs=[]
59 with torch.no_grad():
60 for _ in range(steps+1): xs.append(x.clone()); x=rk4(oscillator,x,dt)
61 return torch.stack(xs,1)
62
63def train_model(model, data, epochs=100, dt=.08):
64 opt=torch.optim.Adam(model.parameters(),lr=3e-3); t0=time.perf_counter()
65 for _ in range(epochs):
66 # random one-step minibatch, same data and objective
67 b=torch.randint(data.shape[0],(min(64,data.shape[0]),)); k=torch.randint(data.shape[1]-1,(len(b),))
68 x=data[b,k]; target=data[b,k+1]
69 pred=rk4(model,x,dt)
70 loss=((pred-target)**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
71 return time.perf_counter()-t0
72
73def rollout(model,x,steps,dt):
74 out=[x.clone()]
75 # PH evaluation still needs input autograd for grad_z H; do not wrap in no_grad.
76 was_training = model.training
77 model.eval()
78 with torch.enable_grad():
79 for _ in range(steps): out.append(rk4(model,out[-1],dt))
80 if was_training: model.train()
81 return torch.stack(out,1)
82
83def main():
84 d=4; z=torch.randn(200,d)
85 # Structural prediction 1: J antisymmetry is exactly zero for every scale.
86 scales=[0.0,.25,.5,1.,2.,4.]
87 skew=[]; psd=[]; drift=[]; predicted=[]
88 ph=PHField(d=d).eval()
89 with torch.no_grad():
90 for s in scales:
91 A=torch.randn(200,d,d)*s; L=torch.randn(200,d,d)*s
92 J=A-A.transpose(1,2); R=L@L.transpose(1,2)+ph.eps*torch.eye(d)
93 skew.append(float((J+J.transpose(1,2)).abs().max()))
94 psd.append(float(torch.linalg.eigvalsh(R).amin()))
95 # use fixed random g; predicted dH/dt has affine lambda dependence
96 g=torch.randn(200,d); val=-(torch.bmm(R,g.unsqueeze(-1)).squeeze(-1)*g).sum(1)
97 drift.append(float(val.mean()))
98 predicted.append(float(-ph.eps*(g*g).sum(1).mean()-s*s*((torch.bmm((torch.randn(200,d,d)*0+L),g.unsqueeze(-1)).squeeze(-1))**2).sum(1).mean()))
99 # Cleaner dissipation sweep: fixed L0,g gives exact prediction slope in lambda^2.
100 L0=torch.randn(200,d,d); g0=torch.randn(200,d); eps=.03; lam=torch.tensor(scales)
101 observed=[]; theory=[]
102 base=eps*(g0*g0).sum(1).mean(); coeff=(torch.bmm(L0,g0.unsqueeze(-1)).squeeze(-1)**2).sum(1).mean()
103 for s in scales:
104 R=s*s*(L0@L0.transpose(1,2))+eps*torch.eye(d)
105 observed.append(float(-(torch.bmm(R,g0.unsqueeze(-1)).squeeze(-1)*g0).sum(1).mean()))
106 theory.append(float(-base-s*s*coeff))
107 # Mini supervised trajectory fit.
108 data=make_data(); ph2=PHField(2,eps=.03); uc=Unconstrained(2)
109 ph_time=train_model(ph2,data); uc_time=train_model(uc,data)
110 test=make_data(n=32,steps=100)[0:]
111 ph_pred=rollout(ph2,test[:,0],100,.08); uc_pred=rollout(uc,test[:,0],100,.08)
112 target=test
113 ph_err=float(torch.sqrt(((ph_pred-target)**2).mean())); uc_err=float(torch.sqrt(((uc_pred-target)**2).mean()))
114 # Unforced energy check over random states and finite RK4 trajectory.
115 H,g,J,R,dh,formula=ph.derivative_terms(z)
116 x=z[:1]; energies=[]
117 with torch.enable_grad():
118 for _ in range(100):
119 energies.append(float(ph.energy(x)[0].detach())); x=rk4(ph,x,.02)
120 energy_increases=sum(np.diff(energies)>1e-10)
121 result={"device":DEVICE,"structural":{"max_skew_residual_by_scale":dict(zip(scales,skew)),"min_R_eigenvalue_by_scale":dict(zip(scales,psd)),"max_energy_formula_residual":float((dh-formula).abs().max()),"energy_derivative_max":float(dh.max()),"continuous_energy_nonpositive":bool(dh.max()<=1e-10),"rk4_energy_increases":int(energy_increases)},"dissipation_sweep":{"lambda":scales,"observed_dHdt":observed,"predicted_dHdt":theory,"max_abs_prediction_error":float(max(abs(a-b) for a,b in zip(observed,theory))),"slope_observed_vs_lambda2":float(np.polyfit(np.array(scales)**2,observed,1)[0]),"slope_predicted":float(-coeff)},"mini_experiment":{"ph_rmse":ph_err,"unconstrained_rmse":uc_err,"ph_train_seconds":ph_time,"unconstrained_train_seconds":uc_time,"steps":100,"ph_params":sum(p.numel() for p in ph2.parameters()),"unconstrained_params":sum(p.numel() for p in uc.parameters())}}
122 print(json.dumps(result,indent=2))
123
124if __name__=='__main__': main()