Port-Hamiltonian Neural ODE / experiment.py

Mechanism confirmed, baseline not beaten

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