Kolmogorov-Lie Unitary Layer / stage2_bench.py

✓✓ Beats tuned baseline

Raw ⬇ ZIP
  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()