import sys, json, random from pathlib import Path import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import get_dataset, make_model, train_model, sweep_baseline, make_report SEEDS = [0,1,2,3,4,5,6,7] GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 1e-2}] EPOCHS = 18 class BraidDynamics(nn.Module): """Set-state-inspired recurrent replacement for rnn_small. Each of the 8 (theta, omega, control) observations is a singleton set; the learned skew connection transports a k-dimensional fiber by Cayley. """ def __init__(self, k=8, hidden=64): super().__init__(); self.k=k; self.hidden=hidden self.obs = nn.Sequential(nn.Linear(3, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh()) self.gen = nn.Linear(32, k*k) self.gate = nn.Linear(32, 1) self.zproj = nn.Linear(32, hidden) self.innov = nn.Sequential(nn.Linear(hidden+k, hidden), nn.Tanh(), nn.Linear(hidden, k)) self.head = nn.Linear(hidden+k, 1) def forward(self, x): b=x.shape[0]; seq=x.view(b,8,3); h=x.new_zeros(b,self.k); z=None eye=torch.eye(self.k,device=x.device,dtype=x.dtype).expand(b,-1,-1) for t in range(8): z=self.obs(seq[:,t]); raw=self.gen(z).view(b,self.k,self.k) A=torch.sigmoid(self.gate(z)).view(b,1,1)*(raw-raw.transpose(1,2)) R=torch.linalg.solve(eye-0.5*A, eye+0.5*A) h=(R@h.unsqueeze(-1)).squeeze(-1) h=h+self.innov(torch.cat((torch.tanh(self.zproj(z)),h),-1)) return self.head(torch.cat((torch.tanh(self.zproj(z)),h),-1)) def signature(self, x): dev=next(self.parameters()).device x=x.to(dev) with torch.no_grad(): b=x.shape[0]; seq=x.view(b,8,3); h=x.new_zeros(b,self.k); dr=[]; nr=[] eye=torch.eye(self.k,device=x.device,dtype=x.dtype).expand(b,-1,-1) for t in range(8): z=self.obs(seq[:,t]); raw=self.gen(z).view(b,self.k,self.k) A=torch.sigmoid(self.gate(z)).view(b,1,1)*(raw-raw.transpose(1,2)) R=torch.linalg.solve(eye-.5*A,eye+.5*A); before=h.norm(dim=1) h=(R@h.unsqueeze(-1)).squeeze(-1); after=h.norm(dim=1) dr.append((A+A.transpose(1,2)).norm(dim=(1,2)).mean().item()) nr.append((after-before).abs().mean().item()) h=h+self.innov(torch.cat((torch.tanh(self.zproj(z)),h),-1)) return {'predicted': 'skew residual near zero and transport-only norm drift near zero', 'observed_mean_skew_residual': float(np.mean(dr)), 'observed_mean_transport_norm_drift': float(np.mean(nr)), 'confirmed': bool(max(dr)<1e-5 and max(nr)<1e-4)} def seed_all(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) def ds_info(ds): return tuple(ds['xtr'].shape[1:]), int(ds['ytr'].shape[1]) if ds['ytr'].ndim>1 else 1 def run_base(cfg): def fn(seed): seed_all(seed); ds=get_dataset('dynamics', seed, n_train=400, n_test=200) net=make_model('rnn_small', tuple(ds['xtr'].shape[1:]), 1) _, metric, _=train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'], batch=128, log=lambda *_:None) return metric return fn def run_idea(cfg, collect=False): nets=[] def fn(seed): seed_all(seed); ds=get_dataset('dynamics', seed, n_train=400, n_test=200) net=BraidDynamics(); trained, metric, _=train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'], batch=128, log=lambda *_:None) if collect: nets.append((trained, ds)) return metric return fn, nets if __name__=='__main__': # Baseline sweep uses the same three learning rates later tried by idea. base=sweep_baseline(run_base, GRID, seeds=SEEDS) from bench.protocol import evaluate idea_sweep=[] for cfg in GRID: r=evaluate(run_idea(cfg)[0], seeds=SEEDS) idea_sweep.append({'cfg':cfg, 'mean':r['mean'], 'per_seed':r['per_seed']}) idea_cfg=min(GRID, key=lambda c: next(z['mean'] for z in idea_sweep if z['cfg']==c)) idea_fn,nets=run_idea(idea_cfg, collect=True) idea=evaluate(idea_fn, seeds=SEEDS) # Signature is computed from trained idea models on their benchmark inputs. sigs=[] for net,ds in nets: if net is not None: sigs.append(net.signature(ds['xte'])) sig={'predicted': 'skew residual near zero and transport-only norm drift near zero', 'observed_mean_skew_residual': float(np.mean([s['observed_mean_skew_residual'] for s in sigs])), 'observed_mean_transport_norm_drift': float(np.mean([s['observed_mean_transport_norm_drift'] for s in sigs])), 'confirmed': bool(sigs and all(s['confirmed'] for s in sigs))} report=make_report('dynamics','rnn_small',base,idea,{'mechanism_signature':sig, 'method_note':'Braid transport replaces the GRU recurrence; same dynamics data, epochs, batch, and lr grid.'}) report['idea']['cfg']=idea_cfg report['idea_sweep']=idea_sweep report['baseline']['protocol_note']='8 paired seeds; sweep grid shared with idea; baseline is canonical rnn_small.' Path('bench_report.json').write_text(json.dumps(report,indent=2)) print(json.dumps(report,indent=2))