Braid-Monodromy Set State / run_bench.py

Mechanism confirmed, baseline not beaten

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