Bilinear Input-Conditioned Koopman Cell / bench_experiment.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys,json,random
 2import numpy as np, torch
 3import torch.nn as nn
 4sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
 5from bench import get_dataset,train_model,sweep_baseline,make_report
 6SEEDS=tuple(range(8)); GRID=[{'lr':x,'epochs':18} for x in (1e-3,3e-3,1e-2)]
 7def seed(s):
 8 random.seed(s); np.random.seed(s); torch.manual_seed(s)
 9 if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
10class Additive(nn.Module):
11 def __init__(self,h=48):
12  super().__init__(); self.inp=nn.Linear(3,h); self.A=nn.Linear(h,h,bias=False); self.D=nn.Linear(3,h,bias=False); self.b=nn.Parameter(torch.zeros(h)); self.head=nn.Linear(h,1)
13 def forward(self,x):
14  q=x.view(x.shape[0],-1,3); z=torch.zeros(x.shape[0],self.A.out_features,device=x.device)
15  for v in q.transpose(0,1): z=torch.tanh(self.A(z)+self.inp(v)+self.D(v)+self.b)
16  return self.head(z)
17class Bilinear(nn.Module):
18 def __init__(self,h=48,r=4):
19  super().__init__(); self.inp=nn.Linear(3,h); self.A=nn.Linear(h,h,bias=False); self.D=nn.Linear(3,h,bias=False); self.U=nn.Parameter(torch.randn(3,h,r)*.03); self.V=nn.Parameter(torch.randn(3,h,r)*.03); self.b=nn.Parameter(torch.zeros(h)); self.head=nn.Linear(h,1)
20 def forward(self,x):
21  q=x.view(x.shape[0],-1,3); z=torch.zeros(x.shape[0],self.A.out_features,device=x.device)
22  for v in q.transpose(0,1):
23   inter=torch.zeros_like(z)
24   for j in range(3): inter=inter+v[:,j:j+1]*(z@self.V[j])@self.U[j].T
25   z=torch.tanh(self.A(z)+self.inp(v)+self.D(v)+inter+self.b)
26  return self.head(z)
27def run(kind,cfg,s,ret=False):
28 seed(s); d=get_dataset('dynamics',s,n_train=400,n_test=120); m=Bilinear() if kind=='idea' else Additive(); n,metric,h=train_model(m,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128); return (metric,n,d) if ret else metric
29def main():
30 def maker(c): return lambda s: run('base',c,s)
31 base=sweep_baseline(maker,GRID,seeds=(0,1,2,3))
32 # Evaluate idea at every shared grid setting; report best on the same full paired seeds.
33 ir=[]
34 for c in GRID:
35  vals=[run('idea',c,s) for s in SEEDS]; ir.append({'cfg':c,'mean':float(np.mean(vals)),'per_seed':vals})
36 best=min(ir,key=lambda x:x['mean']); idea={'mean':best['mean'],'std':float(np.std(best['per_seed'])),'per_seed':best['per_seed'],'n':8,'best_cfg':best['cfg'],'grid':ir}
37 # trained-model signature: action-conditioned prediction change versus action-swapped prediction.
38 sig=[]
39 for s in SEEDS:
40  cfg=best['cfg']; metric,m,d=run('idea',cfg,s,True); m.eval(); x=d['xte'].clone(); xa=x.clone(); xa.view(-1,8,3)[:,:,2]*=-1
41  dev=next(m.parameters()).device; x=x.to(dev); xa=xa.to(dev)
42  with torch.no_grad(): p=m(x).cpu().numpy().ravel(); pa=m(xa).cpu().numpy().ravel()
43  y=d['yte'].numpy().ravel(); sig.append([float(np.mean(np.abs(p-pa))),float(np.mean(np.abs(p-y)))])
44 sig=np.asarray(sig); signature={'observed_action_flip_effect':float(sig[:,0].mean()),'prediction_error':float(sig[:,1].mean()),'ratio':float(sig[:,0].mean()/(sig[:,1].mean()+1e-12)),'confirmed':bool(sig[:,0].mean()>1e-4)}
45 rep=make_report('dynamics','custom_bilinear_vs_additive',base,idea,{'mechanism_signature':signature,'architecture':'matched recurrent latent transition; only bilinear interaction differs','parameter_note':'rank=4, hidden=48'})
46 Path='bench_report.json'; open(Path,'w').write(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
47if __name__=='__main__': main()