Global-Local Koopman Latent Dynamics / run_experiment.py

Mechanism failed

Raw ⬇ ZIP
  1import json, math, time, random
  2from pathlib import Path
  3import numpy as np
  4import torch
  5from torch import nn
  6
  7SEED = 3130
  8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
  9try:
 10    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
 11    if device.type == 'cuda':
 12        torch.cuda.empty_cache()
 13except Exception:
 14    device = torch.device('cpu')
 15
 16# ---------- Stage 1: direct numerical verification of the stated recurrence ----------
 17def verify_math():
 18    rng = np.random.default_rng(SEED)
 19    dg, dl, H = 3, 6, 20
 20    # Stable block operators, as recommended by the idea.
 21    Ag = np.array([[.97,-.12,.0],[.12,.97,.0],[0,0,.93]])
 22    Al = np.zeros((dl, dl))
 23    for i, (a,b) in enumerate([(0.985, .08), (0.96,.18), (0.92,.25)]):
 24        Al[2*i:2*i+2,2*i:2*i+2] = [[a,-b],[b,a]]
 25    B = rng.normal(0,.04,(dg+dl,2))
 26    z0 = rng.normal(size=dg+dl); us = rng.normal(size=(H,2))
 27    A = np.zeros((dg+dl,dg+dl)); A[:dg,:dg]=Ag; A[dg:,dg:]=Al
 28    zs = [z0]
 29    for u in us: zs.append(A @ zs[-1] + B @ u)
 30    zs = np.asarray(zs)
 31    # Same update computed as separate channels; this tests the block formula itself.
 32    zg, zl = z0[:dg].copy(), z0[dg:].copy(); separate=[np.r_[zg,zl]]
 33    for u in us:
 34        zg = Ag @ zg + B[:dg] @ u
 35        zl = Al @ zl + B[dg:] @ u
 36        separate.append(np.r_[zg,zl])
 37    recurrence_error = float(np.max(np.abs(zs-np.asarray(separate))))
 38    # Stability signal: with no controls, the norm contracts because both blocks do.
 39    no_control = [z0.copy()]
 40    for _ in range(H): no_control.append(A @ no_control[-1])
 41    norms = np.linalg.norm(no_control, axis=1)
 42    contraction_ratio = float(norms[-1]/norms[0])
 43    spectral_global = float(max(abs(np.linalg.eigvals(Ag))))
 44    spectral_local = float(max(abs(np.linalg.eigvals(Al))))
 45    return dict(recurrence_max_abs_error=recurrence_error,
 46                spectral_radius_global=spectral_global,
 47                spectral_radius_local=spectral_local,
 48                zero_input_norm_ratio=contraction_ratio,
 49                stability_observed=(spectral_global < 1 and spectral_local < 1 and contraction_ratio < 1))
 50
 51# ---------- Tiny nonlinear observation system with known global/local latent structure ----------
 52def make_data(n=1100, T=24):
 53    rng = np.random.default_rng(SEED+1)
 54    dg, dl = 3, 6
 55    Ag = np.array([[.975,-.11,.0],[.11,.975,.0],[0,0,.94]], dtype=np.float32)
 56    Al = np.zeros((dl,dl), dtype=np.float32)
 57    for i,(a,b) in enumerate([(0.985,.075),(.965,.16),(.93,.22)]):
 58        Al[2*i:2*i+2,2*i:2*i+2] = [[a,-b],[b,a]]
 59    # Mild global-to-local physical coupling makes the observation nonlinear while
 60    # preserving a useful approximately block-structured latent representation.
 61    xs = np.zeros((n,T,12), np.float32)
 62    for k in range(n):
 63        g = rng.normal(0,.8,dg).astype(np.float32)
 64        l = rng.normal(0,.8,dl).astype(np.float32)
 65        for t in range(T):
 66            # nonlinear observation: global coordinates, local coordinates, and products
 67            prod = (g[0] * l[::2] + g[1] * l[1::2]).astype(np.float32)
 68            xs[k,t] = np.r_[g, l, prod]
 69            g = Ag @ g + rng.normal(0,.012,dg).astype(np.float32)
 70            l = Al @ l + np.repeat(np.tanh(g[:3]),2)[:dl].astype(np.float32)*.018
 71    return torch.tensor(xs)
 72
 73class MLP(nn.Module):
 74    def __init__(self, a,b,c): super().__init__(); self.net=nn.Sequential(nn.Linear(a,b),nn.Tanh(),nn.Linear(b,c))
 75    def forward(self,x): return self.net(x)
 76
 77class BlockKoopman(nn.Module):
 78    def __init__(self):
 79        super().__init__(); self.eg=MLP(12,32,3); self.el=MLP(12,32,6)
 80        self.Ag=nn.Parameter(torch.eye(3)*.96); self.Al=nn.Parameter(torch.eye(6)*.96)
 81        self.dec=MLP(9,40,12)
 82    def forward(self,x0,H):
 83        g,l=self.eg(x0),self.el(x0); out=[]
 84        for _ in range(H):
 85            g=torch.einsum('bi,ji->bj',g,self.Ag); l=torch.einsum('bi,ji->bj',l,self.Al)
 86            out.append(self.dec(torch.cat([g,l],1)))
 87        return torch.stack(out,1)
 88
 89class GRUBaseline(nn.Module):
 90    def __init__(self):
 91        super().__init__(); self.enc=MLP(12,32,12); self.gru=nn.GRUCell(12,12); self.dec=MLP(12,40,12)
 92    def forward(self,x0,H):
 93        h=self.enc(x0); out=[]
 94        for _ in range(H):
 95            h=self.gru(torch.zeros_like(h),h); out.append(self.dec(h))
 96        return torch.stack(out,1)
 97
 98def train_eval(model, train, test, epochs=75):
 99    model.to(device); opt=torch.optim.Adam(model.parameters(),lr=3e-3)
100    xtr=train[:,0].to(device); ytr=train[:,1:11].to(device)
101    bs=64; t0=time.perf_counter(); model.train()
102    for ep in range(epochs):
103        perm=torch.randperm(len(xtr),device=device)
104        for ii in range(0,len(xtr),bs):
105            ix=perm[ii:ii+bs]; pred=model(xtr[ix],10)
106            loss=((pred-ytr[ix])**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
107    if device.type=='cuda': torch.cuda.synchronize()
108    train_seconds=time.perf_counter()-t0
109    model.eval();
110    with torch.no_grad():
111        pred=model(test[:,0].to(device),20).cpu(); target=test[:,1:21]
112    mse=((pred-target)**2).mean(2).mean(0).numpy()
113    return dict(mse_1=float(mse[0]),mse_10=float(mse[9]),mse_20=float(mse[19]),train_seconds=train_seconds,
114                params=sum(p.numel() for p in model.parameters())), model
115
116def main():
117    math_check=verify_math()
118    data=make_data(); train=data[:900]; test=data[900:]
119    # Identical initialization order and data; each model gets the same objective/horizon.
120    torch.manual_seed(SEED+10); base, _=train_eval(GRUBaseline(),train,test)
121    torch.manual_seed(SEED+10); idea, model=train_eval(BlockKoopman(),train,test)
122    with torch.no_grad():
123        rho_g=max(abs(np.linalg.eigvals(model.Ag.detach().cpu().numpy())))
124        rho_l=max(abs(np.linalg.eigvals(model.Al.detach().cpu().numpy())))
125    idea['learned_rho_global']=float(rho_g); idea['learned_rho_local']=float(rho_l)
126    result={'device':str(device),'math_check':math_check,'baseline':base,'idea':idea,
127            'lower_mse_at_20':idea['mse_20'] < base['mse_20']}
128    Path('results.json').write_text(json.dumps(result,indent=2))
129    print(json.dumps(result,indent=2))
130
131if __name__=='__main__': main()