Adaptive Physics-Lifted Koopman State Space / stage2_bench.py
Beats tuned baseline
1import sys, json, math, random
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, evaluate, make_report
8
9SEEDS = list(range(8))
10GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 1e-2}]
11EPOCHS, NTRAIN, NTEST = 12, 600, 300
12
13class LiftedKoopman(nn.Module):
14 """A learned observable z=[x,x^2], linear controlled transition z'=Az+Bu."""
15 def __init__(self, hidden=64, horizon=8, radius=.97):
16 super().__init__()
17 self.hidden, self.horizon, self.radius = hidden, horizon, radius
18 self.enc = nn.Linear(3, hidden)
19 self.A = nn.Parameter(torch.eye(2*hidden) * .85)
20 self.B = nn.Parameter(torch.zeros(2*hidden, 3))
21 self.dec = nn.Linear(2*hidden, 1)
22 nn.init.normal_(self.B, std=.015)
23 nn.init.normal_(self.dec.weight, std=.03)
24 nn.init.zeros_(self.dec.bias)
25
26 def stable_A(self):
27 # Differentiable scalar projection keeps the learned transition bounded.
28 a = self.A
29 scale = torch.linalg.matrix_norm(a, ord=2).clamp_min(1e-6)
30 return a * torch.minimum(torch.ones_like(scale), self.radius / scale)
31
32 def forward(self, x):
33 seq = x.view(x.shape[0], -1, 3)
34 h = torch.tanh(self.enc(seq[:, 0]))
35 z = torch.cat((h, h*h), dim=1)
36 A = self.stable_A()
37 for k in range(1, seq.shape[1]):
38 z = z @ A.T + seq[:, k] @ self.B.T
39 # one-step forecast from the final observed state/action
40 z = z @ A.T + seq[:, -1] @ self.B.T
41 return self.dec(z)
42
43
44def seed_all(seed):
45 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
46 if torch.cuda.is_available():
47 try: torch.cuda.manual_seed_all(seed)
48 except Exception: pass
49
50def baseline_fn(cfg):
51 def run(seed):
52 seed_all(seed); d=get_dataset('dynamics', seed, NTRAIN, NTEST)
53 net=make_model('rnn_small', d['input_shape'], d['out_dim'])
54 _, metric, _=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *a,**k:None)
55 return metric
56 return run
57
58def idea_fn(cfg):
59 def run(seed):
60 seed_all(seed); d=get_dataset('dynamics', seed, NTRAIN, NTEST)
61 net=LiftedKoopman(); _, metric, _=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *a,**k:None)
62 return metric
63 return run
64
65def rls_check():
66 rng=np.random.default_rng(17); d,n,lam=4,35,.93
67 th=np.zeros((2,d)); P=np.eye(d)*4
68 X=[]; Y=[]
69 for _ in range(n):
70 r=rng.normal(size=d); y=np.array([.4*r[0]-.2*r[1]+.1*r[2], .2*r[3]])
71 pr=P@r; pn=(P-np.outer(pr,pr)/(lam+r@pr))/lam
72 th=th+np.outer(y-th@r, r@pn); P=pn; X.append(r); Y.append(y)
73 X=np.asarray(X); Y=np.asarray(Y); w=lam**np.arange(n-1,-1,-1)
74 batch=(Y.T*w)@[email protected](X.T@(w[:,None]*X)+(lam**n)*np.eye(d)/4)
75 return float(np.max(np.abs(th-batch))), 1/(1-lam)
76
77def signature():
78 seed=SEEDS[0]; seed_all(seed); d=get_dataset('dynamics',seed,NTRAIN,NTEST)
79 net=LiftedKoopman(); net,_,_=train_model(net,d,epochs=EPOCHS,lr=3e-3,batch=128,log=lambda *a,**k:None)
80 with torch.no_grad():
81 A=net.stable_A().detach().cpu().numpy(); rho=float(np.max(np.abs(np.linalg.eigvals(A))))
82 dev=next(net.parameters()).device
83 x0=d['xte'][:1].to(dev).view(1,8,3)[:,0]
84 h=torch.tanh(net.enc(x0))
85 z=torch.cat((h, h**2),1)
86 norms=[]
87 for k in range(9):
88 norms.append(float(torch.linalg.vector_norm(z).cpu())); z=z@net.stable_A().T
89 ratios=[norms[k+1]/max(norms[k],1e-12) for k in range(8)]
90 observed=float(np.mean(ratios[-3:]))
91 return {'rho_trained_A':rho,'predicted_asymptotic_norm_ratio':rho,'observed_mean_last3_norm_ratio':observed,'horizon_norms':norms,'confirmed':abs(observed-rho)<.08,'source':'trained benchmark Koopman model'}
92
93def main():
94 err,tau=rls_check()
95 base=sweep_baseline(baseline_fn,GRID,seeds=[0,1,2])
96 # Shared union parity: evaluate all three idea settings on all eight paired seeds.
97 idea_runs=[]
98 for cfg in GRID:
99 r=evaluate(idea_fn(cfg),seeds=SEEDS)
100 idea_runs.append({'cfg':cfg,'result':r})
101 best=min(idea_runs,key=lambda q:q['result']['mean'])
102 rep=make_report('dynamics','rnn_small',base,best['result'],{
103 'mechanism_signature':signature(),
104 'math_check':{'rls_batch_max_abs_error':err,'predicted_adaptation_timescale':tau},
105 'idea_sweep':idea_runs,
106 'protocol_note':'Baseline sweep uses same lr union; final comparison is eight paired seeds.'})
107 Path('bench_report.json').write_text(json.dumps(rep,indent=2))
108 print(json.dumps(rep,indent=2))
109if __name__=='__main__': main()