Braid-Monodromy Set State / run_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json, random
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, sweep_baseline, make_report
8
9SEEDS = [0,1,2,3,4,5,6,7]
10GRID = [{'lr': 1e-3}, {'lr': 3e-3}, {'lr': 1e-2}]
11EPOCHS = 18
12
13class BraidDynamics(nn.Module):
14 """Set-state-inspired recurrent replacement for rnn_small.
15 Each of the 8 (theta, omega, control) observations is a singleton set;
16 the learned skew connection transports a k-dimensional fiber by Cayley.
17 """
18 def __init__(self, k=8, hidden=64):
19 super().__init__(); self.k=k; self.hidden=hidden
20 self.obs = nn.Sequential(nn.Linear(3, 32), nn.Tanh(), nn.Linear(32, 32), nn.Tanh())
21 self.gen = nn.Linear(32, k*k)
22 self.gate = nn.Linear(32, 1)
23 self.zproj = nn.Linear(32, hidden)
24 self.innov = nn.Sequential(nn.Linear(hidden+k, hidden), nn.Tanh(), nn.Linear(hidden, k))
25 self.head = nn.Linear(hidden+k, 1)
26 def forward(self, x):
27 b=x.shape[0]; seq=x.view(b,8,3); h=x.new_zeros(b,self.k); z=None
28 eye=torch.eye(self.k,device=x.device,dtype=x.dtype).expand(b,-1,-1)
29 for t in range(8):
30 z=self.obs(seq[:,t]); raw=self.gen(z).view(b,self.k,self.k)
31 A=torch.sigmoid(self.gate(z)).view(b,1,1)*(raw-raw.transpose(1,2))
32 R=torch.linalg.solve(eye-0.5*A, eye+0.5*A)
33 h=(R@h.unsqueeze(-1)).squeeze(-1)
34 h=h+self.innov(torch.cat((torch.tanh(self.zproj(z)),h),-1))
35 return self.head(torch.cat((torch.tanh(self.zproj(z)),h),-1))
36 def signature(self, x):
37 dev=next(self.parameters()).device
38 x=x.to(dev)
39 with torch.no_grad():
40 b=x.shape[0]; seq=x.view(b,8,3); h=x.new_zeros(b,self.k); dr=[]; nr=[]
41 eye=torch.eye(self.k,device=x.device,dtype=x.dtype).expand(b,-1,-1)
42 for t in range(8):
43 z=self.obs(seq[:,t]); raw=self.gen(z).view(b,self.k,self.k)
44 A=torch.sigmoid(self.gate(z)).view(b,1,1)*(raw-raw.transpose(1,2))
45 R=torch.linalg.solve(eye-.5*A,eye+.5*A); before=h.norm(dim=1)
46 h=(R@h.unsqueeze(-1)).squeeze(-1); after=h.norm(dim=1)
47 dr.append((A+A.transpose(1,2)).norm(dim=(1,2)).mean().item())
48 nr.append((after-before).abs().mean().item())
49 h=h+self.innov(torch.cat((torch.tanh(self.zproj(z)),h),-1))
50 return {'predicted': 'skew residual near zero and transport-only norm drift near zero',
51 'observed_mean_skew_residual': float(np.mean(dr)),
52 'observed_mean_transport_norm_drift': float(np.mean(nr)),
53 'confirmed': bool(max(dr)<1e-5 and max(nr)<1e-4)}
54
55def seed_all(s):
56 random.seed(s); np.random.seed(s); torch.manual_seed(s)
57
58def ds_info(ds):
59 return tuple(ds['xtr'].shape[1:]), int(ds['ytr'].shape[1]) if ds['ytr'].ndim>1 else 1
60
61def run_base(cfg):
62 def fn(seed):
63 seed_all(seed); ds=get_dataset('dynamics', seed, n_train=400, n_test=200)
64 net=make_model('rnn_small', tuple(ds['xtr'].shape[1:]), 1)
65 _, metric, _=train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'], batch=128, log=lambda *_:None)
66 return metric
67 return fn
68
69def run_idea(cfg, collect=False):
70 nets=[]
71 def fn(seed):
72 seed_all(seed); ds=get_dataset('dynamics', seed, n_train=400, n_test=200)
73 net=BraidDynamics(); trained, metric, _=train_model(net, ds, epochs=EPOCHS, lr=cfg['lr'], batch=128, log=lambda *_:None)
74 if collect: nets.append((trained, ds))
75 return metric
76 return fn, nets
77
78if __name__=='__main__':
79 # Baseline sweep uses the same three learning rates later tried by idea.
80 base=sweep_baseline(run_base, GRID, seeds=SEEDS)
81 from bench.protocol import evaluate
82 idea_sweep=[]
83 for cfg in GRID:
84 r=evaluate(run_idea(cfg)[0], seeds=SEEDS)
85 idea_sweep.append({'cfg':cfg, 'mean':r['mean'], 'per_seed':r['per_seed']})
86 idea_cfg=min(GRID, key=lambda c: next(z['mean'] for z in idea_sweep if z['cfg']==c))
87 idea_fn,nets=run_idea(idea_cfg, collect=True)
88 idea=evaluate(idea_fn, seeds=SEEDS)
89 # Signature is computed from trained idea models on their benchmark inputs.
90 sigs=[]
91 for net,ds in nets:
92 if net is not None: sigs.append(net.signature(ds['xte']))
93 sig={'predicted': 'skew residual near zero and transport-only norm drift near zero',
94 'observed_mean_skew_residual': float(np.mean([s['observed_mean_skew_residual'] for s in sigs])),
95 'observed_mean_transport_norm_drift': float(np.mean([s['observed_mean_transport_norm_drift'] for s in sigs])),
96 'confirmed': bool(sigs and all(s['confirmed'] for s in sigs))}
97 report=make_report('dynamics','rnn_small',base,idea,{'mechanism_signature':sig,
98 'method_note':'Braid transport replaces the GRU recurrence; same dynamics data, epochs, batch, and lr grid.'})
99 report['idea']['cfg']=idea_cfg
100 report['idea_sweep']=idea_sweep
101 report['baseline']['protocol_note']='8 paired seeds; sweep grid shared with idea; baseline is canonical rnn_small.'
102 Path('bench_report.json').write_text(json.dumps(report,indent=2))
103 print(json.dumps(report,indent=2))