import sys, json, random, itertools 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, train_model, sweep_baseline, make_report SEEDS = tuple(range(8)) LR_GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 1e-2}] EPOCHS = 10 NTR, NTE = 400, 100 H = 16 class DenseRNN(nn.Module): """Standard tanh recurrent controller, used as the matched baseline.""" def __init__(self, out_dim=1): super().__init__() self.inp = nn.Linear(3, H) self.rec = nn.Linear(H, H, bias=False) self.head = nn.Linear(H, out_dim) def forward(self, x): z = x.view(x.shape[0], -1, 3) h = torch.zeros(x.shape[0], H, device=x.device) for t in range(z.shape[1]): h = torch.tanh(self.inp(z[:, t]) + self.rec(h)) return self.head(h) def recurrent_norms(self, x): # Trained-model behavior of the dense recurrent map on benchmark states. with torch.no_grad(): z=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],H,device=x.device) vals=[] for t in range(z.shape[1]): vals.append(torch.linalg.vector_norm(self.rec(h),dim=1).mean()) h=torch.tanh(self.inp(z[:,t])+self.rec(h)) return torch.stack(vals) class KLU_RNN(nn.Module): """Same recurrent shell, replacing W h by an input-conditioned Lie product U(x)h.""" def __init__(self, out_dim=1, K=4): super().__init__(); self.K=K self.inp=nn.Linear(3,H); self.head=nn.Linear(H,out_dim) self.raw_skew=nn.Parameter(torch.randn(K,H,H)*0.03) self.psi=nn.ModuleList([nn.ModuleList([nn.Sequential(nn.Linear(1,6),nn.Tanh(),nn.Linear(6,1)) for _ in range(3)]) for _ in range(K)]) self.coef=nn.ModuleList([nn.Sequential(nn.Linear(1,8),nn.Tanh(),nn.Linear(8,1),nn.Tanh()) for _ in range(K)]) def unitary(self, x): B=x.shape[0]; eye=torch.eye(H,device=x.device,dtype=x.dtype).expand(B,H,H) U=eye for k in range(self.K): s=sum(self.psi[k][j](x[:,j:j+1]) for j in range(3)) a=self.coef[k](s).view(B,1,1) q=(self.raw_skew[k]-self.raw_skew[k].T)*0.5 U=torch.bmm(torch.matrix_exp(a*q.unsqueeze(0)),U) return U def forward(self,x): z=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],H,device=x.device) for t in range(z.shape[1]): h=torch.tanh(self.inp(z[:,t])+torch.bmm(self.unitary(z[:,t]),h.unsqueeze(2)).squeeze(2)) return self.head(h) def recurrent_norms(self,x): with torch.no_grad(): z=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],H,device=x.device); vals=[]; residual=[] eye=torch.eye(H,device=x.device) for t in range(z.shape[1]): U=self.unitary(z[:,t]); uh=torch.bmm(U,h.unsqueeze(2)).squeeze(2) vals.append(torch.linalg.vector_norm(uh,dim=1).mean()) residual.append((U.transpose(1,2)@U-eye).norm(dim=(1,2)).mean()) h=torch.tanh(self.inp(z[:,t])+uh) return torch.stack(vals), torch.stack(residual) def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) def run(kind, cfg, seed, capture=False): seed_all(seed) ds=get_dataset('dynamics', seed, n_train=NTR, n_test=NTE) model=(DenseRNN(ds['out_dim']) if kind=='baseline' else KLU_RNN(ds['out_dim'])).float() net, metric, hist=train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=128, log=lambda *_: None) if net is None: return float('nan'), {} extra={} if capture: dev=next(net.parameters()).device; xx=ds['xte'].to(dev) if kind=='baseline': v=net.recurrent_norms(xx); extra={'recurrent_norm_mean':float(v.mean()),'recurrent_norm_abs_drift':float((v- v[:,0] if v.ndim>1 else v-v[0]).abs().mean())} else: v,r=net.recurrent_norms(xx); extra={'recurrent_norm_mean':float(v.mean()),'recurrent_norm_abs_drift':float((v-v[0]).abs().mean()),'unitarity_residual':float(r.mean())} return float(metric), extra def make_train(kind,cfg): return lambda seed: run(kind,cfg,seed)[0] def main(): # Independent cheap numerical check, before benchmark training. torch.manual_seed(388); A=torch.randn(H,H); Q=(A-A.T)*.15 E=torch.matrix_exp(Q); math_res=float((E.T@E-torch.eye(H)).norm()) base=sweep_baseline(lambda c: make_train('baseline',c), LR_GRID, seeds=SEEDS) idea_trials=[] for c in LR_GRID: r={"cfg":c,"result":__import__('bench').evaluate(make_train('idea',c),SEEDS)} idea_trials.append(r) best=min(idea_trials,key=lambda q:q['result']['mean']) idea=best['result']; best_cfg=best['cfg'] sigvals=[] for s in SEEDS: _, e=run('idea',best_cfg,s,True); _, b=run('baseline',base['best_cfg'],s,True) sigvals.append({'seed':s,'idea':e,'baseline':b}) i_norm=np.mean([q['idea']['recurrent_norm_mean'] for q in sigvals]); b_norm=np.mean([q['baseline']['recurrent_norm_mean'] for q in sigvals]) i_res=np.mean([q['idea']['unitarity_residual'] for q in sigvals]) signature={'prediction':'KLU recurrent map preserves the norm of each incoming hidden state while dense recurrence need not','observed_idea_recurrent_norm':float(i_norm),'observed_baseline_recurrent_output_norm':float(b_norm),'observed_idea_unitarity_residual':float(i_res),'confirmed':bool(i_res < 1e-4)} report=make_report('dynamics','rnn_small',base,idea,extra={'mechanism_signature':signature,'baseline_grid':LR_GRID,'idea_grid':idea_trials,'math_check_exponential_residual':math_res,'protocol_note':'dynamics selected because the idea claims stability/control through norm-preserving recurrent transformations'}) report['idea']['best_cfg']=best_cfg Path('bench_report.json').write_text(json.dumps(report,indent=2)) print(json.dumps(report,indent=2)) if __name__=='__main__': main()