import sys, json, math, random from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, sweep_baseline, evaluate, make_report SEEDS = list(range(8)) GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 1e-2}] EPOCHS, NTRAIN, NTEST = 12, 600, 300 class LiftedKoopman(nn.Module): """A learned observable z=[x,x^2], linear controlled transition z'=Az+Bu.""" def __init__(self, hidden=64, horizon=8, radius=.97): super().__init__() self.hidden, self.horizon, self.radius = hidden, horizon, radius self.enc = nn.Linear(3, hidden) self.A = nn.Parameter(torch.eye(2*hidden) * .85) self.B = nn.Parameter(torch.zeros(2*hidden, 3)) self.dec = nn.Linear(2*hidden, 1) nn.init.normal_(self.B, std=.015) nn.init.normal_(self.dec.weight, std=.03) nn.init.zeros_(self.dec.bias) def stable_A(self): # Differentiable scalar projection keeps the learned transition bounded. a = self.A scale = torch.linalg.matrix_norm(a, ord=2).clamp_min(1e-6) return a * torch.minimum(torch.ones_like(scale), self.radius / scale) def forward(self, x): seq = x.view(x.shape[0], -1, 3) h = torch.tanh(self.enc(seq[:, 0])) z = torch.cat((h, h*h), dim=1) A = self.stable_A() for k in range(1, seq.shape[1]): z = z @ A.T + seq[:, k] @ self.B.T # one-step forecast from the final observed state/action z = z @ A.T + seq[:, -1] @ self.B.T return self.dec(z) def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): try: torch.cuda.manual_seed_all(seed) except Exception: pass def baseline_fn(cfg): def run(seed): seed_all(seed); d=get_dataset('dynamics', seed, NTRAIN, NTEST) net=make_model('rnn_small', d['input_shape'], d['out_dim']) _, metric, _=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *a,**k:None) return metric return run def idea_fn(cfg): def run(seed): seed_all(seed); d=get_dataset('dynamics', seed, NTRAIN, NTEST) net=LiftedKoopman(); _, metric, _=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *a,**k:None) return metric return run def rls_check(): rng=np.random.default_rng(17); d,n,lam=4,35,.93 th=np.zeros((2,d)); P=np.eye(d)*4 X=[]; Y=[] for _ in range(n): r=rng.normal(size=d); y=np.array([.4*r[0]-.2*r[1]+.1*r[2], .2*r[3]]) pr=P@r; pn=(P-np.outer(pr,pr)/(lam+r@pr))/lam th=th+np.outer(y-th@r, r@pn); P=pn; X.append(r); Y.append(y) X=np.asarray(X); Y=np.asarray(Y); w=lam**np.arange(n-1,-1,-1) batch=(Y.T*w)@X@np.linalg.inv(X.T@(w[:,None]*X)+(lam**n)*np.eye(d)/4) return float(np.max(np.abs(th-batch))), 1/(1-lam) def signature(): seed=SEEDS[0]; seed_all(seed); d=get_dataset('dynamics',seed,NTRAIN,NTEST) net=LiftedKoopman(); net,_,_=train_model(net,d,epochs=EPOCHS,lr=3e-3,batch=128,log=lambda *a,**k:None) with torch.no_grad(): A=net.stable_A().detach().cpu().numpy(); rho=float(np.max(np.abs(np.linalg.eigvals(A)))) dev=next(net.parameters()).device x0=d['xte'][:1].to(dev).view(1,8,3)[:,0] h=torch.tanh(net.enc(x0)) z=torch.cat((h, h**2),1) norms=[] for k in range(9): norms.append(float(torch.linalg.vector_norm(z).cpu())); z=z@net.stable_A().T ratios=[norms[k+1]/max(norms[k],1e-12) for k in range(8)] observed=float(np.mean(ratios[-3:])) 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'} def main(): err,tau=rls_check() base=sweep_baseline(baseline_fn,GRID,seeds=[0,1,2]) # Shared union parity: evaluate all three idea settings on all eight paired seeds. idea_runs=[] for cfg in GRID: r=evaluate(idea_fn(cfg),seeds=SEEDS) idea_runs.append({'cfg':cfg,'result':r}) best=min(idea_runs,key=lambda q:q['result']['mean']) rep=make_report('dynamics','rnn_small',base,best['result'],{ 'mechanism_signature':signature(), 'math_check':{'rls_batch_max_abs_error':err,'predicted_adaptation_timescale':tau}, 'idea_sweep':idea_runs, 'protocol_note':'Baseline sweep uses same lr union; final comparison is eight paired seeds.'}) Path('bench_report.json').write_text(json.dumps(rep,indent=2)) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()