import json, math, random import numpy as np import torch from torch import nn SEED = 2045 random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED) torch.set_default_dtype(torch.float64) DEVICE = 'cuda' if torch.cuda.is_available() else 'cpu' try: if DEVICE == 'cuda': torch.zeros(1, device='cuda') except Exception: DEVICE = 'cpu' # Exact scalar certificate sanity check. def scalar_mu(a, lam, k, q, u=0.0): M = math.exp(k*u) # f=a*x, A=a, M(u)=exp(k*u), u_dot=q. mdot = k*q*M return (mdot + 2*a*M + 2*lam*M) / M def verify_math(): a, lam, k = -0.8, 0.2, 1.5 qstar = -2*(a + lam)/k qs = np.linspace(qstar-1.0, qstar+1.0, 101) mus = np.array([scalar_mu(a, lam, k, float(q)) for q in qs]) slope, intercept = np.polyfit(qs, mus, 1) # Mechanism predictions: boundary qstar, linear q sensitivity k, and no sensitivity k=0. q_observed = -intercept/slope qgrid = np.linspace(-3, 3, 25) flat = np.array([scalar_mu(a, lam, 0.0, float(q)) for q in qgrid]) return { 'predicted_boundary_q': float(qstar), 'observed_boundary_q': float(q_observed), 'predicted_dmu_dq': float(k), 'observed_dmu_dq': float(slope), 'k0_max_abs_q_variation': float(np.max(np.abs(flat-flat[0]))), 'frozen_metric_mu_q1': float(scalar_mu(a,lam,k,0.0)), 'total_derivative_mu_q1': float(scalar_mu(a,lam,k,1.0)), 'predicted_sign_q_minus': float(scalar_mu(a,lam,k,qstar-0.5)), 'predicted_sign_q_plus': float(scalar_mu(a,lam,k,qstar+0.5)), } class Field(nn.Module): def __init__(self): super().__init__() self.net = nn.Sequential(nn.Linear(3,32), nn.Tanh(), nn.Linear(32,32), nn.Tanh(), nn.Linear(32,2)) def forward(self,x,u): return self.net(torch.cat([x,u],-1)) class Metric(nn.Module): def __init__(self): super().__init__() self.net=nn.Sequential(nn.Linear(3,24),nn.Tanh(),nn.Linear(24,3)) def forward(self,x,u): z=self.net(torch.cat([x,u],-1)) L=torch.zeros(x.shape[0],2,2,device=x.device,dtype=x.dtype) L[:,0,0]=torch.nn.functional.softplus(z[:,0])+0.05 L[:,1,0]=z[:,1]; L[:,1,1]=torch.nn.functional.softplus(z[:,2])+0.05 return L@L.transpose(-1,-2)+1e-3*torch.eye(2,device=x.device,dtype=x.dtype) def true_f(x,u): return torch.stack([x[:,1], (1-x[:,0]**2)*x[:,1]-x[:,0]+u[:,0]],1) def metric_directional_derivative(metric, x, u, xdot, udot): # Per-sample total derivative dM/dt = dM/dx xdot + dM/du udot. M=metric(x,u) dm=torch.zeros_like(M) for r in range(2): for c in range(2): gx=torch.autograd.grad(M[:,r,c].sum(),x,create_graph=True,retain_graph=True)[0] gu=torch.autograd.grad(M[:,r,c].sum(),u,create_graph=True,retain_graph=True)[0] dm[:,r,c]=(gx*xdot).sum(-1)+(gu*udot).sum(-1) return dm def jacobian_batch(y,z): rows=[] for i in range(y.shape[1]): rows.append(torch.autograd.grad(y[:,i].sum(),z,create_graph=True,retain_graph=True)[0]) return torch.stack(rows,1) def train(kind, steps=350): dev=torch.device(DEVICE); model=Field().to(dev); metric=Metric().to(dev) if kind=='adaptive' else None opt=torch.optim.Adam(list(model.parameters())+([] if metric is None else list(metric.parameters())),lr=2e-3) t=torch.linspace(0,1,65,device=dev); dt=t[1]-t[0] x=torch.zeros(48,2,device=dev); x[:,0]=torch.randn(48,device=dev)*.7; x[:,1]=torch.randn(48,device=dev)*.3 u=torch.sin(2.0*t)[None,:,None].repeat(48,1,1) target=[] with torch.no_grad(): z=x for j in range(64): z=z+dt*true_f(z,u[:,j]); target.append(z) target=torch.stack(target,1) for step in range(steps): j=np.random.randint(0,64); xx=target[:,j].detach().clone().requires_grad_(True); uu=u[:,j].detach().clone().requires_grad_(True) pred=model(xx,uu); loss=((pred-true_f(xx,uu))**2).mean() if kind!='unconstrained': A=jacobian_batch(pred,xx) if kind=='fixed': M=torch.eye(2,device=dev).expand(xx.shape[0],2,2) else: M=metric(xx,uu) S=A.transpose(-1,-2)@M+M@A+0.4*M if kind=='adaptive': udot=2*torch.cos(2*t[j]).expand_as(uu) S=S+metric_directional_derivative(metric,xx,uu,pred,udot) # robust-free certificate; softplus of largest eigenvalue. ev=torch.linalg.eigvalsh(torch.linalg.solve(M,S)).real[:,-1] loss=loss+0.15*torch.nn.functional.softplus(ev).mean() opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(list(model.parameters()),5); opt.step() # Evaluate total derivative certificate on held-out rapidly varying controls. xx=torch.randn(256,2,device=dev).requires_grad_(True); te=torch.rand(256,1,device=dev) uu=torch.sin(8*te).requires_grad_(True); udot=8*torch.cos(8*te) pred=model(xx,uu); A=jacobian_batch(pred,xx) if kind=='unconstrained' or kind=='fixed': M=torch.eye(2,device=dev).expand(256,2,2) else: M=metric(xx,uu) S=A.transpose(-1,-2)@M+M@A+0.4*M if kind=='adaptive': S=S+metric_directional_derivative(metric,xx,uu,pred,udot) mu=torch.linalg.eigvalsh(torch.linalg.solve(M,S)).real[:,-1].detach().cpu().numpy() return {'mean_mu':float(mu.mean()),'max_mu':float(mu.max()),'violation_fraction':float((mu>0).mean())} def main(): math_check=verify_math(); results={k:train(k) for k in ['unconstrained','fixed','adaptive']} out={'device':DEVICE,'math_check':math_check,'models':results} with open('results.json','w') as f: json.dump(out,f,indent=2) print(json.dumps(out,indent=2)) if __name__=='__main__': main()