Koopman-MPC Trust Region for Neural Rollouts / koopman_mpc_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
  1import json, math, os, random
  2import numpy as np
  3import torch
  4from torch import nn
  5
  6SEED=7
  7np.random.seed(SEED); random.seed(SEED); torch.manual_seed(SEED)
  8try:
  9    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 10except Exception:
 11    device=torch.device('cpu')
 12
 13def pend_step(x, dt=.08):
 14    # normalized damped pendulum, x=[angle, angular velocity]
 15    th, om=x[...,0], x[...,1]
 16    return np.stack([th+dt*om, om+dt*(-np.sin(th)-.12*om)], axis=-1)
 17
 18def make_data(ntraj=180, length=35):
 19    xs=[]
 20    for _ in range(ntraj):
 21        x=np.array([np.random.uniform(-2.5,2.5), np.random.uniform(-2.0,2.0)])
 22        for _ in range(length):
 23            y=pend_step(x); xs.append((x.copy(),y.copy())); x=y
 24    return np.asarray([a for a,b in xs],np.float32), np.asarray([b for a,b in xs],np.float32)
 25
 26class MLP(nn.Module):
 27    def __init__(self):
 28        super().__init__(); self.net=nn.Sequential(nn.Linear(2,32),nn.Tanh(),nn.Linear(32,32),nn.Tanh(),nn.Linear(32,2))
 29    def forward(self,x): return self.net(x)
 30
 31def cap_spectral(A,r):
 32    rho=max(abs(np.linalg.eigvals(A)))
 33    return A*(r/rho) if rho>r else A.copy(), float(rho)
 34
 35def prediction_mats(A,B,N):
 36    n,m=B.shape; AA=np.zeros((N*n,n)); BB=np.zeros((N*n,N*m))
 37    # Stack Z=[z_0,...,z_{N-1}].  Then z_k=A^k z_0+
 38    # sum_{j<k} A^(k-1-j) B v_j.
 39    p=np.eye(n)
 40    for k in range(N):
 41        AA[k*n:(k+1)*n]=p
 42        for j in range(k):
 43            BB[k*n:(k+1)*n,j*m:(j+1)*m]=(np.linalg.matrix_power(A,k-1-j)@B)
 44        p=p@A
 45    return AA,BB
 46
 47def mpc_correction(z, neural_next, A, N=8, vmax=.20, delta=1.0):
 48    # B=I; Q tracks the stable Koopman prediction, R regularizes correction.
 49    n=2; B=np.eye(n); AA,BB=prediction_mats(A,B,N)
 50    target=AA@z
 51    # The neural proposal is the initial state of the controlled rollout;
 52    # optimize corrections toward the adapted Koopman reference trajectory.
 53    # This avoids the degenerate zero correction obtained by targeting AA z
 54    # from the same AA z initial condition.
 55    base=np.tile(neural_next, N)
 56    H=BB.T@BB+.08*np.eye(N*n); g=BB.T@(base-target)
 57    try: V=np.linalg.solve(H,-g)
 58    except np.linalg.LinAlgError: V=np.zeros(N*n)
 59    v=np.clip(V[:n],-vmax,vmax)
 60    raw=neural_next+v
 61    # trust region around current observed latent state; radial projection is the
 62    # small-QP feasibility projection for this two-dimensional toy.
 63    d=raw-z; norm=np.linalg.norm(d)
 64    if norm>delta: raw=z+d*(delta/norm)
 65    return raw, float(np.linalg.norm(v))
 66
 67def main():
 68    X,Y=make_data(); tx=torch.tensor(X,device=device); ty=torch.tensor(Y,device=device)
 69    model=MLP().to(device); opt=torch.optim.Adam(model.parameters(),lr=2e-3)
 70    for epoch in range(260):
 71        p=model(tx); loss=((p-ty)**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
 72    with torch.no_grad(): train_mse=float(((model(tx)-ty)**2).mean().cpu())
 73    # least-squares adapted linear latent dynamics
 74    A=np.linalg.lstsq(X,Y,rcond=None)[0].T
 75    A_cap,rho=cap_spectral(A,.95)
 76    # verify condensed formula against iterative dynamics
 77    B=np.eye(2); z=np.array([.3,-.4]); V=np.random.randn(6*2)*.03
 78    AA,BB=prediction_mats(A_cap,B,6); zs=[]; q=z.copy()
 79    # The standard matrices stack z_0,...,z_{N-1}; the first row is z_0.
 80    for k in range(6):
 81        zs.append(q.copy())
 82        q=A_cap@q+B@V[k*2:(k+1)*2]
 83    mat_err=float(np.max(np.abs(np.asarray(zs).reshape(-1)-(AA@z+BB@V))))
 84    # spectral growth check, same initial norm and 100 steps
 85    z0=np.array([1.,1.]); u=z0.copy(); c=z0.copy()
 86    for _ in range(100): u=A@u; c=A_cap@c
 87    growth_un=float(np.linalg.norm(u)/np.linalg.norm(z0)); growth_cap=float(np.linalg.norm(c)/np.linalg.norm(z0))
 88    # long rollout errors from held-out initial conditions
 89    starts=[]
 90    for _ in range(24): starts.append(np.array([np.random.uniform(-2.5,2.5),np.random.uniform(-2,2)],np.float32))
 91    horizons=60; errs_base=[]; errs_mpc=[]; norms_base=[]; norms_mpc=[]
 92    for s in starts:
 93        truth=s.copy(); zb=s.copy(); zm=s.copy()
 94        for _ in range(horizons):
 95            truth=pend_step(truth)
 96            with torch.no_grad(): nnnext=model(torch.tensor(zb,device=device,dtype=torch.float32)).cpu().numpy()
 97            zb=nnnext
 98            with torch.no_grad(): nnnext2=model(torch.tensor(zm,device=device,dtype=torch.float32)).cpu().numpy()
 99            zm,_=mpc_correction(zm,nnnext2,A_cap,N=8,vmax=.20,delta=.55)
100            errs_base.append(np.linalg.norm(zb-truth)); errs_mpc.append(np.linalg.norm(zm-truth))
101            norms_base.append(np.linalg.norm(zb)); norms_mpc.append(np.linalg.norm(zm))
102    # trust-radius sweep exposes bias vs extrapolation in this implementation
103    sweep=[]
104    for delta in [.15,.3,.55,1.0,2.0]:
105        ee=[]
106        for s in starts[:12]:
107            truth=s.copy(); z=s.copy()
108            for _ in range(40):
109                truth=pend_step(truth)
110                with torch.no_grad(): nnn=model(torch.tensor(z,device=device,dtype=torch.float32)).cpu().numpy()
111                z,_=mpc_correction(z,nnn,A_cap,delta=delta)
112                ee.append(np.linalg.norm(z-truth))
113        sweep.append([delta,float(np.mean(ee))])
114    out={'device':str(device),'train_one_step_mse':train_mse,'rho_unconstrained':rho,
115         'condensed_matrix_max_error':mat_err,'growth_100_unconstrained':growth_un,
116         'growth_100_capped':growth_cap,'baseline_mean_60step_error':float(np.mean(errs_base)),
117         'idea_mean_60step_error':float(np.mean(errs_mpc)),'baseline_mean_pred_norm':float(np.mean(norms_base)),
118         'idea_mean_pred_norm':float(np.mean(norms_mpc)),'trust_radius_sweep':sweep,
119         'config':{'n_train_pairs':len(X),'n_test':len(starts),'horizon':60,'spectral_cap':.95,'vmax':.20}}
120    open('results.json','w').write(json.dumps(out,indent=2))
121    print(json.dumps(out,indent=2))
122if __name__=='__main__': main()