Support-Sparse Koopman World Model / stage2_bench.py
Failed on benchmark
1import sys, json, random, time
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
8
9OUT='bench_report.json'; SEEDS=tuple(range(8)); LRS=[1e-3,3e-3,6e-3]
10EPOCHS=4; BATCH=128; NTR=400; NTE=200
11
12def seed_all(s):
13 random.seed(s); np.random.seed(s); torch.manual_seed(s)
14 if torch.cuda.is_available():
15 try: torch.cuda.manual_seed_all(s)
16 except Exception: pass
17
18def baseline_one(cfg, seed, return_model=False):
19 seed_all(seed+10000); ds=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
20 m=make_model('rnn_small',ds['input_shape'],1)
21 m,metric,hist=train_model(m,ds,epochs=EPOCHS,lr=cfg['lr'],batch=BATCH,weight_decay=cfg.get('wd',0.0),log=lambda *_:None)
22 return (metric,m,ds) if return_model else metric
23
24class SparseKoopRNN(nn.Module):
25 def __init__(self,latent=16):
26 super().__init__(); self.rnn=nn.GRU(3,64,batch_first=True); self.head=nn.Linear(64,1)
27 self.enc=nn.Sequential(nn.Linear(3,32),nn.Tanh(),nn.Linear(32,latent))
28 self.dec=nn.Sequential(nn.Linear(latent,32),nn.Tanh(),nn.Linear(32,3))
29 self.K=nn.Parameter(torch.eye(latent)+.01*torch.randn(latent,latent))
30 def forward(self,x):
31 _,h=self.rnn(x.view(x.shape[0],-1,3)); return self.head(h[-1])
32 def aux(self,x,H=6):
33 states=x.view(x.shape[0],-1,3); zs=self.enc(states); rec=((self.dec(zs)-states)**2).mean()
34 pred=zs[:,0]; lat=roll=support=0.; eps=1e-6
35 for h in range(1,H+1):
36 pred=pred@self.K.T; lat+=((pred-zs[:,h])**2).mean(); roll+=((self.dec(pred)-states[:,h])**2).mean()
37 support+=(pred.abs()/(pred.abs()+eps)-zs[:,h-1].abs()/(zs[:,h-1].abs()+eps)).abs().mean()
38 return rec,lat/H,roll/H,support/H
39
40def train_idea_device(cfg,seed,ds,dev):
41 seed_all(seed+20000); m=SparseKoopRNN(16).to(dev); x,y=ds['xtr'].to(dev),ds['ytr'].to(dev)
42 opt=torch.optim.Adam(m.parameters(),lr=cfg['lr'])
43 for ep in range(EPOCHS):
44 m.train(); p=torch.randperm(len(x),device=dev)
45 for i in range(0,len(x),BATCH):
46 xb=x[p[i:i+BATCH]]; yb=y[p[i:i+BATCH]]; task=((m(xb)-yb)**2).mean()
47 rec,lat,roll,sup=m.aux(xb)
48 loss=task+.15*rec+.35*lat+.35*roll+cfg['lam']*m.enc(xb.view(xb.shape[0],-1,3)[:,0]).abs().mean()+cfg['eta']*sup
49 opt.zero_grad(); loss.backward(); torch.nn.utils.clip_grad_norm_(m.parameters(),10); opt.step()
50 m.eval()
51 with torch.no_grad(): metric=float(((m(ds['xte'].to(dev))-ds['yte'].to(dev))**2).mean().cpu())
52 return metric,m,dev
53
54def idea_one(cfg,seed,return_model=False):
55 ds=get_dataset('dynamics',seed,n_train=NTR,n_test=NTE)
56 dev=torch.device('cpu')
57 if torch.cuda.is_available():
58 try: torch.zeros(1,device='cuda'); dev=torch.device('cuda')
59 except Exception: pass
60 try: result=train_idea_device(cfg,seed,ds,dev)
61 except RuntimeError:
62 result=train_idea_device(cfg,seed,ds,torch.device('cpu'))
63 return (result[0], result[1], ds, result[2]) if return_model else result[0]
64
65def aggregate(vals):
66 return {'mean':float(np.mean(vals)),'std':float(np.std(vals)),'per_seed':[float(v) for v in vals],'n':len(vals)}
67
68def main():
69 t=time.time(); grid=[{'lr':lr,'wd':0.0} for lr in LRS]
70 sw=sweep_baseline(lambda c: lambda s: baseline_one(c,s),grid,seeds=SEEDS)
71 best_cfg=sw['best_cfg']; base={'best_cfg':best_cfg,'sweep':sw['sweep'],'full':sw['full']}
72 idea_grid=[{'lr':best_cfg['lr'],'lam':5e-4,'eta':2e-3},{'lr':1e-3,'lam':5e-4,'eta':2e-3},{'lr':6e-3,'lam':5e-4,'eta':2e-3}]
73 runs=[{'cfg':c,'result':aggregate([idea_one(c,s) for s in SEEDS])} for c in idea_grid]
74 best=min(runs,key=lambda z:z['result']['mean']); metric,m,ds,dev=idea_one(best['cfg'],0,True)
75 with torch.no_grad():
76 st=ds['xte'].to(dev).view(-1,8,3); z=m.enc(st); pred=z[:,0]; lat=[]; ro=[]; sup=[]
77 for h in range(1,7):
78 pred=pred@m.K.T; lat.append(((pred-z[:,h])**2).mean().item()); ro.append(((m.dec(pred)-st[:,h])**2).mean().item())
79 sup.append((pred.abs()/(pred.abs()+1e-6)-z[:,h-1].abs()/(z[:,h-1].abs()+1e-6)).abs().mean().item())
80 active=(z.abs()>0.25).float().mean().item()
81 sig={'one_step_latent_mse_mean':float(np.mean(lat)),'decoded_rollout_mse_h6':float(ro[-1]),'support_change_mean':float(np.mean(sup)),'active_fraction_tau_025':active,'confirmed':False}
82 rep=make_report('dynamics','rnn_small',base,best['result'],extra=sig); rep['idea_sweep']=runs; rep['runtime_sec']=time.time()-t
83 Path(OUT).write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
84if __name__=='__main__': main()