Input-Aware Contracting Neural ODE / experiment.py
Failed on benchmark
1import json, math, random
2import numpy as np
3import torch
4from torch import nn
5
6SEED = 2045
7random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
8torch.set_default_dtype(torch.float64)
9DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu'
10try:
11 if DEVICE == 'cuda': torch.zeros(1, device='cuda')
12except Exception:
13 DEVICE = 'cpu'
14
15# Exact scalar certificate sanity check.
16def scalar_mu(a, lam, k, q, u=0.0):
17 M = math.exp(k*u)
18 # f=a*x, A=a, M(u)=exp(k*u), u_dot=q.
19 mdot = k*q*M
20 return (mdot + 2*a*M + 2*lam*M) / M
21
22def verify_math():
23 a, lam, k = -0.8, 0.2, 1.5
24 qstar = -2*(a + lam)/k
25 qs = np.linspace(qstar-1.0, qstar+1.0, 101)
26 mus = np.array([scalar_mu(a, lam, k, float(q)) for q in qs])
27 slope, intercept = np.polyfit(qs, mus, 1)
28 # Mechanism predictions: boundary qstar, linear q sensitivity k, and no sensitivity k=0.
29 q_observed = -intercept/slope
30 qgrid = np.linspace(-3, 3, 25)
31 flat = np.array([scalar_mu(a, lam, 0.0, float(q)) for q in qgrid])
32 return {
33 'predicted_boundary_q': float(qstar),
34 'observed_boundary_q': float(q_observed),
35 'predicted_dmu_dq': float(k),
36 'observed_dmu_dq': float(slope),
37 'k0_max_abs_q_variation': float(np.max(np.abs(flat-flat[0]))),
38 'frozen_metric_mu_q1': float(scalar_mu(a,lam,k,0.0)),
39 'total_derivative_mu_q1': float(scalar_mu(a,lam,k,1.0)),
40 'predicted_sign_q_minus': float(scalar_mu(a,lam,k,qstar-0.5)),
41 'predicted_sign_q_plus': float(scalar_mu(a,lam,k,qstar+0.5)),
42 }
43
44class Field(nn.Module):
45 def __init__(self):
46 super().__init__()
47 self.net = nn.Sequential(nn.Linear(3,32), nn.Tanh(), nn.Linear(32,32), nn.Tanh(), nn.Linear(32,2))
48 def forward(self,x,u): return self.net(torch.cat([x,u],-1))
49
50class Metric(nn.Module):
51 def __init__(self):
52 super().__init__()
53 self.net=nn.Sequential(nn.Linear(3,24),nn.Tanh(),nn.Linear(24,3))
54 def forward(self,x,u):
55 z=self.net(torch.cat([x,u],-1))
56 L=torch.zeros(x.shape[0],2,2,device=x.device,dtype=x.dtype)
57 L[:,0,0]=torch.nn.functional.softplus(z[:,0])+0.05
58 L[:,1,0]=z[:,1]; L[:,1,1]=torch.nn.functional.softplus(z[:,2])+0.05
59 return L@L.transpose(-1,-2)+1e-3*torch.eye(2,device=x.device,dtype=x.dtype)
60
61def true_f(x,u):
62 return torch.stack([x[:,1], (1-x[:,0]**2)*x[:,1]-x[:,0]+u[:,0]],1)
63
64def metric_directional_derivative(metric, x, u, xdot, udot):
65 # Per-sample total derivative dM/dt = dM/dx xdot + dM/du udot.
66 M=metric(x,u)
67 dm=torch.zeros_like(M)
68 for r in range(2):
69 for c in range(2):
70 gx=torch.autograd.grad(M[:,r,c].sum(),x,create_graph=True,retain_graph=True)[0]
71 gu=torch.autograd.grad(M[:,r,c].sum(),u,create_graph=True,retain_graph=True)[0]
72 dm[:,r,c]=(gx*xdot).sum(-1)+(gu*udot).sum(-1)
73 return dm
74
75def jacobian_batch(y,z):
76 rows=[]
77 for i in range(y.shape[1]):
78 rows.append(torch.autograd.grad(y[:,i].sum(),z,create_graph=True,retain_graph=True)[0])
79 return torch.stack(rows,1)
80
81def train(kind, steps=350):
82 dev=torch.device(DEVICE); model=Field().to(dev); metric=Metric().to(dev) if kind=='adaptive' else None
83 opt=torch.optim.Adam(list(model.parameters())+([] if metric is None else list(metric.parameters())),lr=2e-3)
84 t=torch.linspace(0,1,65,device=dev); dt=t[1]-t[0]
85 x=torch.zeros(48,2,device=dev); x[:,0]=torch.randn(48,device=dev)*.7; x[:,1]=torch.randn(48,device=dev)*.3
86 u=torch.sin(2.0*t)[None,:,None].repeat(48,1,1)
87 target=[]
88 with torch.no_grad():
89 z=x
90 for j in range(64):
91 z=z+dt*true_f(z,u[:,j]); target.append(z)
92 target=torch.stack(target,1)
93 for step in range(steps):
94 j=np.random.randint(0,64); xx=target[:,j].detach().clone().requires_grad_(True); uu=u[:,j].detach().clone().requires_grad_(True)
95 pred=model(xx,uu); loss=((pred-true_f(xx,uu))**2).mean()
96 if kind!='unconstrained':
97 A=jacobian_batch(pred,xx)
98 if kind=='fixed': M=torch.eye(2,device=dev).expand(xx.shape[0],2,2)
99 else:
100 M=metric(xx,uu)
101 S=A.transpose(-1,-2)@M+M@A+0.4*M
102 if kind=='adaptive':
103 udot=2*torch.cos(2*t[j]).expand_as(uu)
104 S=S+metric_directional_derivative(metric,xx,uu,pred,udot)
105 # robust-free certificate; softplus of largest eigenvalue.
106 ev=torch.linalg.eigvalsh(torch.linalg.solve(M,S)).real[:,-1]
107 loss=loss+0.15*torch.nn.functional.softplus(ev).mean()
108 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(list(model.parameters()),5); opt.step()
109 # Evaluate total derivative certificate on held-out rapidly varying controls.
110 xx=torch.randn(256,2,device=dev).requires_grad_(True); te=torch.rand(256,1,device=dev)
111 uu=torch.sin(8*te).requires_grad_(True); udot=8*torch.cos(8*te)
112 pred=model(xx,uu); A=jacobian_batch(pred,xx)
113 if kind=='unconstrained' or kind=='fixed': M=torch.eye(2,device=dev).expand(256,2,2)
114 else: M=metric(xx,uu)
115 S=A.transpose(-1,-2)@M+M@A+0.4*M
116 if kind=='adaptive': S=S+metric_directional_derivative(metric,xx,uu,pred,udot)
117 mu=torch.linalg.eigvalsh(torch.linalg.solve(M,S)).real[:,-1].detach().cpu().numpy()
118 return {'mean_mu':float(mu.mean()),'max_mu':float(mu.max()),'violation_fraction':float((mu>0).mean())}
119
120def main():
121 math_check=verify_math(); results={k:train(k) for k in ['unconstrained','fixed','adaptive']}
122 out={'device':DEVICE,'math_check':math_check,'models':results}
123 with open('results.json','w') as f: json.dump(out,f,indent=2)
124 print(json.dumps(out,indent=2))
125if __name__=='__main__': main()