Support-Sparse Koopman World Model / experiment.py

Failed on benchmark

Raw ⬇ ZIP
 1import json, math, random, time
 2from pathlib import Path
 3import numpy as np
 4import torch
 5from torch import nn
 6
 7SEED=2805
 8random.seed(SEED); np.random.seed(SEED); torch.manual_seed(SEED)
 9torch.set_num_threads(min(8, torch.get_num_threads()))
10try:
11    device=torch.device('cuda' if torch.cuda.is_available() else 'cpu')
12    if device.type=='cuda': torch.zeros(1, device=device)
13except Exception:
14    device=torch.device('cpu')
15
16def math_checks():
17    eps=1e-6
18    # Prediction 1: away from zero, ds/da = eps/(a+eps)^2.
19    amps=np.array([1e-4,1e-3,1e-2,1e-1,1.0])
20    delta=1e-8
21    s=lambda x: abs(x)/(abs(x)+eps)
22    observed=[(s(a+delta)-s(a))/delta for a in amps]
23    predicted=[eps/(a+eps)**2 for a in amps]
24    rel=float(np.max(np.abs(np.array(observed)-predicted)/(np.array(predicted)+1e-30)))
25    # Prediction 2: for K=rI, rollout norm is exactly r^h times initial norm.
26    rs=np.array([0.7,0.9,1.0,1.1,1.3]); h=8; v=np.ones(4)/2
27    rollout_obs=[np.linalg.norm((r**h)*v) for r in rs]
28    rollout_pred=[(r**h)*np.linalg.norm(v) for r in rs]
29    rollout_err=float(np.max(np.abs(np.array(rollout_obs)-rollout_pred)))
30    # Prediction 3: for |delta| << a, two-sided support change is
31    # 2 eps |delta|/(a+eps)^2. Use a delta that never crosses zero.
32    delta2=1e-8
33    support_obs=[]; support_pred=[]
34    for a in amps:
35        obs=abs(s(a+delta2)-s(a))+abs(s(a-delta2)-s(a))
36        pr=2*eps/(a+eps)**2*delta2
37        support_obs.append(obs); support_pred.append(pr)
38    support_rel=float(np.max(np.abs(np.array(support_obs)-support_pred)/(np.array(support_pred)+1e-30)))
39    return {'support_derivative':{'amplitudes':amps.tolist(),'observed':observed,'predicted':predicted,'max_relative_error':rel},
40            'rollout_scaling':{'r':rs.tolist(),'h':h,'observed_norm':rollout_obs,'predicted_norm':rollout_pred,'max_abs_error':rollout_err},
41            'support_perturbation':{'amplitudes':amps.tolist(),'delta':delta2,'observed':support_obs,'predicted':support_pred,'max_relative_error':support_rel}}
42
43def make_data(ntraj=180,T=70):
44    allx=[]; labels=[]
45    for basin in [-1,1]:
46        for _ in range(ntraj//2):
47            q=basin+np.random.normal(0,.12); v=np.random.normal(0,.12); traj=[]
48            for t in range(T):
49                traj.append([q,v]); acc=q-q**3-.32*v; q=q+.08*v; v=v+.08*acc
50            allx.append(np.asarray(traj,np.float32)); labels.append(basin)
51    return np.asarray(allx),np.asarray(labels)
52
53class Koopman(nn.Module):
54    def __init__(self,dz=16):
55        super().__init__(); self.enc=nn.Sequential(nn.Linear(2,32),nn.Tanh(),nn.Linear(32,dz)); self.dec=nn.Sequential(nn.Linear(dz,32),nn.Tanh(),nn.Linear(32,2)); self.K=nn.Parameter(torch.eye(dz)+.01*torch.randn(dz,dz))
56    def forward(self,x): return self.enc(x)
57
58def train_model(x,sparse=True,eta=.002,steps=700,H=5,dz=16):
59    torch.manual_seed(SEED+(1 if sparse else 2)+int(eta*100000)); model=Koopman(dz).to(device); opt=torch.optim.Adam(model.parameters(),lr=2e-3); xt=torch.tensor(x,device=device); n,T,_=x.shape; losses=[]
60    for step in range(steps):
61        ids=torch.randint(0,n,(64,),device=device); ts=torch.randint(0,T-H,(64,),device=device); win=torch.stack([xt[ids,j] for j in range(H+1)],1); z0=model(win[:,0]); pred=z0; rec=((model.dec(z0)-win[:,0])**2).mean(); lat=roll=sup=0.; prev_z=z0
62        for j in range(1,H+1):
63            pred=pred@model.K.T; zj=model(win[:,j]); lat+=((pred-zj)**2).mean(); roll+=((model.dec(pred)-win[:,j])**2).mean(); ss=lambda z:z.abs()/(z.abs()+1e-6); sup+=(ss(zj)-ss(prev_z)).abs().mean(); prev_z=zj
64        loss=rec+.5*lat/H+roll/H+(.0005*z0.abs().mean() if sparse else 0)+(eta*sup if sparse else 0); opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(model.parameters(),10); opt.step()
65        if step%100==0: losses.append(float(loss.detach().cpu()))
66    return model,losses
67
68@torch.no_grad()
69def evaluate(model,x,labels,H=20,tau=None):
70    xt=torch.tensor(x,device=device); z=model(xt[:,0]); pred=z; outs=[]
71    for h in range(1,H+1): pred=pred@model.K.T; outs.append(model.dec(pred))
72    mse=float(((torch.stack(outs,1)-xt[:,1:H+1])**2).mean().cpu()); zz=model(xt.reshape(-1,2)).reshape(len(x),len(x[0]),-1)
73    if tau is None: tau=float(torch.quantile(zz.abs().flatten(),.75).cpu())
74    masks=(zz[:,0].abs()>tau).cpu().numpy().astype(np.int8); keys=[tuple(a.tolist()) for a in masks]; uniq={k:i for i,k in enumerate(sorted(set(keys)))}; pg=np.array([uniq[k]%2 for k in keys]); y=(labels>0).astype(int); acc=max(np.mean(pg==y),np.mean(1-pg==y)); ss=zz.abs()/(zz.abs()+1e-6)
75    return {'mse20':mse,'tau':tau,'active_fraction':float(masks.mean()),'mask_basin_acc':float(acc),'support_step_change':float((ss[:,1:]-ss[:,:-1]).abs().mean().cpu())}
76
77def main():
78    t0=time.time(); checks=math_checks(); train_x,train_y=make_data(180,70); test_x,test_y=make_data(60,70); results={}
79    base,_=train_model(train_x,False,eta=0,steps=700); results['dense']=evaluate(base,test_x,test_y)
80    sparse,_=train_model(train_x,True,eta=.002,steps=700); results['sparse']=evaluate(sparse,test_x,test_y)
81    ablation={}
82    for eta in [0.,.001,.005]:
83        m,_=train_model(train_x,True,eta=eta,steps=450); ablation[str(eta)]=evaluate(m,test_x,test_y)
84    report={'seed':SEED,'device':str(device),'math_checks':checks,'baseline_vs_idea':results,'eta_ablation':ablation,'runtime_sec':time.time()-t0,'predictions':['support derivative scales as epsilon/(a+epsilon)^2','linear rollout norm scales exactly as spectral radius^h','small support perturbations scale as 2 epsilon |delta|/(a+epsilon)^2']}
85    Path('results.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
86if __name__=='__main__': main()