Input-Aware Contracting Neural ODE / experiment.py

Failed on benchmark

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