Kolmogorov-Lie Unitary Layer / stage2_bench.py
Beats tuned baseline
1import sys, json, random, itertools
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import get_dataset, train_model, sweep_baseline, make_report
9
10SEEDS = tuple(range(8))
11LR_GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 1e-2}]
12EPOCHS = 10
13NTR, NTE = 400, 100
14H = 16
15
16class DenseRNN(nn.Module):
17 """Standard tanh recurrent controller, used as the matched baseline."""
18 def __init__(self, out_dim=1):
19 super().__init__()
20 self.inp = nn.Linear(3, H)
21 self.rec = nn.Linear(H, H, bias=False)
22 self.head = nn.Linear(H, out_dim)
23 def forward(self, x):
24 z = x.view(x.shape[0], -1, 3)
25 h = torch.zeros(x.shape[0], H, device=x.device)
26 for t in range(z.shape[1]):
27 h = torch.tanh(self.inp(z[:, t]) + self.rec(h))
28 return self.head(h)
29 def recurrent_norms(self, x):
30 # Trained-model behavior of the dense recurrent map on benchmark states.
31 with torch.no_grad():
32 z=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],H,device=x.device)
33 vals=[]
34 for t in range(z.shape[1]):
35 vals.append(torch.linalg.vector_norm(self.rec(h),dim=1).mean())
36 h=torch.tanh(self.inp(z[:,t])+self.rec(h))
37 return torch.stack(vals)
38
39class KLU_RNN(nn.Module):
40 """Same recurrent shell, replacing W h by an input-conditioned Lie product U(x)h."""
41 def __init__(self, out_dim=1, K=4):
42 super().__init__(); self.K=K
43 self.inp=nn.Linear(3,H); self.head=nn.Linear(H,out_dim)
44 self.raw_skew=nn.Parameter(torch.randn(K,H,H)*0.03)
45 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)])
46 self.coef=nn.ModuleList([nn.Sequential(nn.Linear(1,8),nn.Tanh(),nn.Linear(8,1),nn.Tanh()) for _ in range(K)])
47 def unitary(self, x):
48 B=x.shape[0]; eye=torch.eye(H,device=x.device,dtype=x.dtype).expand(B,H,H)
49 U=eye
50 for k in range(self.K):
51 s=sum(self.psi[k][j](x[:,j:j+1]) for j in range(3))
52 a=self.coef[k](s).view(B,1,1)
53 q=(self.raw_skew[k]-self.raw_skew[k].T)*0.5
54 U=torch.bmm(torch.matrix_exp(a*q.unsqueeze(0)),U)
55 return U
56 def forward(self,x):
57 z=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],H,device=x.device)
58 for t in range(z.shape[1]):
59 h=torch.tanh(self.inp(z[:,t])+torch.bmm(self.unitary(z[:,t]),h.unsqueeze(2)).squeeze(2))
60 return self.head(h)
61 def recurrent_norms(self,x):
62 with torch.no_grad():
63 z=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],H,device=x.device); vals=[]; residual=[]
64 eye=torch.eye(H,device=x.device)
65 for t in range(z.shape[1]):
66 U=self.unitary(z[:,t]); uh=torch.bmm(U,h.unsqueeze(2)).squeeze(2)
67 vals.append(torch.linalg.vector_norm(uh,dim=1).mean())
68 residual.append((U.transpose(1,2)@U-eye).norm(dim=(1,2)).mean())
69 h=torch.tanh(self.inp(z[:,t])+uh)
70 return torch.stack(vals), torch.stack(residual)
71
72def seed_all(s):
73 random.seed(s); np.random.seed(s); torch.manual_seed(s)
74
75def run(kind, cfg, seed, capture=False):
76 seed_all(seed)
77 ds=get_dataset('dynamics', seed, n_train=NTR, n_test=NTE)
78 model=(DenseRNN(ds['out_dim']) if kind=='baseline' else KLU_RNN(ds['out_dim'])).float()
79 net, metric, hist=train_model(model, ds, epochs=EPOCHS, lr=cfg['lr'], batch=128, log=lambda *_: None)
80 if net is None: return float('nan'), {}
81 extra={}
82 if capture:
83 dev=next(net.parameters()).device; xx=ds['xte'].to(dev)
84 if kind=='baseline':
85 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())}
86 else:
87 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())}
88 return float(metric), extra
89
90def make_train(kind,cfg):
91 return lambda seed: run(kind,cfg,seed)[0]
92
93def main():
94 # Independent cheap numerical check, before benchmark training.
95 torch.manual_seed(388); A=torch.randn(H,H); Q=(A-A.T)*.15
96 E=torch.matrix_exp(Q); math_res=float((E.T@E-torch.eye(H)).norm())
97 base=sweep_baseline(lambda c: make_train('baseline',c), LR_GRID, seeds=SEEDS)
98 idea_trials=[]
99 for c in LR_GRID:
100 r={"cfg":c,"result":__import__('bench').evaluate(make_train('idea',c),SEEDS)}
101 idea_trials.append(r)
102 best=min(idea_trials,key=lambda q:q['result']['mean'])
103 idea=best['result']; best_cfg=best['cfg']
104 sigvals=[]
105 for s in SEEDS:
106 _, e=run('idea',best_cfg,s,True); _, b=run('baseline',base['best_cfg'],s,True)
107 sigvals.append({'seed':s,'idea':e,'baseline':b})
108 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])
109 i_res=np.mean([q['idea']['unitarity_residual'] for q in sigvals])
110 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)}
111 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'})
112 report['idea']['best_cfg']=best_cfg
113 Path('bench_report.json').write_text(json.dumps(report,indent=2))
114 print(json.dumps(report,indent=2))
115if __name__=='__main__': main()