Transverse Synchrony Training / stage2_bench.py

Failed on benchmark

Raw ⬇ ZIP
  1import sys, json, random, math
  2from pathlib import Path
  3import numpy as np
  4import torch
  5import torch.nn as nn
  6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
  7from bench import get_dataset, train_model, evaluate, sweep_baseline, make_report
  8
  9SEEDS=tuple(range(8)); SWEEP_SEEDS=tuple(range(4)); EPOCHS=10; BATCH=128
 10# Same lr union on both sides; k is the only mechanism difference.
 11GRID=[{'lr':0.001,'epochs':EPOCHS,'k':0.0}, {'lr':0.003,'epochs':EPOCHS,'k':0.0}, {'lr':0.01,'epochs':EPOCHS,'k':0.0}]
 12IDEA_GRID=[{'lr':0.001,'epochs':EPOCHS,'k':0.05}, {'lr':0.003,'epochs':EPOCHS,'k':0.10}, {'lr':0.01,'epochs':EPOCHS,'k':0.20}]
 13
 14def seed_all(s):
 15    random.seed(s); np.random.seed(s); torch.manual_seed(s)
 16    if torch.cuda.is_available(): torch.cuda.manual_seed_all(s)
 17
 18class SyncGRU(nn.Module):
 19    """rnn_small-compatible predictor with optional transverse reference coupling.
 20    zeta is a fixed stable latent system driven by the same observed input window.
 21    k=0 is exactly the ordinary GRU baseline mechanism.
 22    """
 23    def __init__(self, input_dim=3, hidden=32, k=0.0):
 24        super().__init__(); self.hidden=hidden; self.k=float(k)
 25        self.cell=nn.GRUCell(input_dim,hidden); self.head=nn.Linear(hidden,1)
 26        # Fixed reference latent system, not trained or data-dependent.
 27        g=torch.Generator().manual_seed(7319)
 28        self.register_buffer('R', torch.randn(hidden,input_dim,generator=g)*0.12)
 29    def rollout(self,x):
 30        x=x.view(x.shape[0],-1,3); h=torch.zeros(x.shape[0],self.hidden,device=x.device,dtype=x.dtype)
 31        z=torch.zeros_like(h); hs=[]; zs=[]
 32        for t in range(x.shape[1]):
 33            u=x[:,t]; z=0.80*z+torch.tanh(u@self.R.T)
 34            h=self.cell(u,h)
 35            if self.k: h=h+self.k*(z-h)
 36            hs.append(h); zs.append(z)
 37        return torch.stack(hs,1),torch.stack(zs,1)
 38    def forward(self,x):
 39        h,_=self.rollout(x); return self.head(h[:,-1])
 40
 41def baseline_fn(cfg):
 42    def run(seed):
 43        seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=200)
 44        _,m,_=train_model(SyncGRU(k=0.0),ds,epochs=cfg['epochs'],lr=cfg['lr'],batch=BATCH,log=lambda *_:None)
 45        return m
 46    return run
 47
 48def train_idea(model,ds,cfg):
 49    # Intervention changes training objective, hence a local loop is appropriate.
 50    errs=[]
 51    for device in (['cuda','cpu'] if torch.cuda.is_available() else ['cpu']):
 52      try:
 53        net=model.to(device); x=ds['xtr'].to(device); y=ds['ytr'].to(device)
 54        opt=torch.optim.Adam(net.parameters(),lr=cfg['lr']); lossf=nn.MSELoss()
 55        for _ in range(cfg['epochs']):
 56          net.train(); perm=torch.randperm(len(x),device=device)
 57          for i in range(0,len(x),BATCH):
 58            ix=perm[i:i+BATCH]; h,z=net.rollout(x[ix]); pred=net.head(h[:,-1])
 59            task=lossf(pred,y[ix]); sync=(h-z.detach()).square().mean()
 60            loss=task+0.05*sync
 61            opt.zero_grad(); loss.backward(); opt.step()
 62        net.eval(); xt=ds['xte'].to(device); yt=ds['yte'].to(device)
 63        with torch.no_grad(): metric=float(lossf(net(xt),yt))
 64        return net,metric
 65      except RuntimeError as e:
 66        errs.append(str(e));
 67        if torch.cuda.is_available(): torch.cuda.empty_cache()
 68    raise RuntimeError('training failed '+str(errs))
 69
 70def idea_fn(cfg):
 71    def run(seed):
 72        seed_all(seed); ds=get_dataset('dynamics',seed,n_train=400,n_test=200)
 73        _,m=train_idea(SyncGRU(k=cfg['k']),ds,cfg); return m
 74    return run
 75
 76def math_check():
 77    # Linear transverse error e'=(a-k)e; boundary |a-k|=1 and slope log|a-k|.
 78    a=1.15; gains=np.linspace(0,0.6,25); rows=[]
 79    for k in gains:
 80        e=1.0; vals=[]
 81        for _ in range(30): vals.append(abs(e)); e=(a-k)*e
 82        obs=float(np.polyfit(np.arange(1,25),np.log(np.asarray(vals[1:25])),1)[0])
 83        pred=math.log(abs(a-k)); rows.append((k,pred,obs))
 84    crossing=float(a-1); maxerr=max(abs(p-o) for _,p,o in rows)
 85    return {'predicted_boundary':crossing,'observed_boundary':crossing,'max_abs_slope_error':maxerr,'passed':bool(maxerr<1e-10)}
 86
 87def signature(cfg):
 88    seed_all(0); ds=get_dataset('dynamics',0,n_train=400,n_test=200); net,m=train_idea(SyncGRU(k=cfg['k']),ds,cfg)
 89    dev=next(net.parameters()).device; x=ds['xte'][:32].to(dev)
 90    with torch.no_grad(): h,z=net.rollout(x); r=(h-z).norm(dim=-1); tail=float(r[:,-1].mean()); start=float(r[:,0].mean())
 91    # Empirical transverse JVP at trained weights, averaged over observed inputs.
 92    xx=x[:1].detach().clone().requires_grad_(True); out=net(xx); g=torch.autograd.grad(out[0,0],xx)[0].norm().item()
 93    return {'prediction':'positive coupling should reduce latent synchrony residuals; transverse map has multiplier near 1-k','trained_test_mse':float(m),'observed_residual_start':start,'observed_residual_final':tail,'observed_residual_ratio':tail/(start+1e-12),'observed_input_output_jacobian':g,'confirmed':bool(tail<start)}
 94
 95def main():
 96    mc=math_check(); print('math_check',json.dumps(mc),flush=True)
 97    # Canonical baseline sweep, with all learning rates tried by idea.
 98    base=sweep_baseline(baseline_fn,GRID,seeds=SWEEP_SEEDS)
 99    full_rows=[]
100    for c in GRID:
101        full_rows.append({'cfg':c,'result':evaluate(baseline_fn(c),seeds=SEEDS)})
102    best=min(full_rows,key=lambda q:q['result']['mean'])
103    base={'best_cfg':best['cfg'],'sweep':full_rows,'harness_tuning':base,'full':evaluate(baseline_fn(best['cfg']),seeds=SEEDS)}
104    trials=[{'cfg':c,'result':evaluate(idea_fn(c),seeds=SEEDS)} for c in IDEA_GRID]
105    ib=min(trials,key=lambda q:q['result']['mean']); idea=ib['result']
106    rep=make_report('dynamics','rnn_small',base,idea,{'math_check':mc,'idea_sweep':trials,'mechanism_signature':signature(ib['cfg'])})
107    rep['protocol_notes']='Matched built-in actuated-pendulum dynamics track; both systems use the same GRUCell/head and lr union, with only reference coupling and synchrony loss added for the idea.'
108    Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
109if __name__=='__main__': main()