Hysteretic Multiscale Sequence Router / bench_hysteretic.py
Failed on benchmark
1import sys,json,random
2import numpy as np, torch
3import torch.nn as nn
4import torch.nn.functional as F
5sys.path.insert(0,'/home/maxwelhelp/all/math2nn')
6from bench import get_dataset,train_model,evaluate,sweep_baseline,make_report
7SEEDS=tuple(range(8)); EPOCHS=10; NTR=400; NTE=200; K=4
8GRID=[{'lr':.0015,'temperature':.7},{'lr':.003,'temperature':1.0},{'lr':.006,'temperature':1.4}]
9class Router(nn.Module):
10 def __init__(self,win,mode,temp=1.,delta=.12,tau=2.):
11 super().__init__(); self.mode=mode; self.temp=temp; self.delta=delta; self.tau=tau
12 d=48; self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.randn(1,win,d)*.02)
13 self.enc=nn.TransformerEncoder(nn.TransformerEncoderLayer(d,2,96,batch_first=True,dropout=0),2)
14 self.score=nn.Linear(d,K); self.experts=nn.ModuleList([nn.Linear(d,1) for _ in range(K)])
15 self.stats={}
16 def forward(self,x):
17 z=self.enc(self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]])
18 s=F.softmax(self.score(z)/self.temp,-1)
19 if self.mode=='soft': g=s; ch=g.argmax(-1)
20 else:
21 # Differentiable leaky hysteresis approximation: cumulative state is detached
22 # for routing decisions, while the soft scores provide the straight-through gradient.
23 b=x.shape[0]; h=torch.full((b,K),1./K,device=x.device); active=torch.zeros(b,dtype=torch.long,device=x.device)
24 outs=[]; choices=[]; switches=0; min_dwell=10**9; last=torch.zeros(b,dtype=torch.long,device=x.device)
25 for t in range(x.shape[1]):
26 h=h+(s[:,t]-h)/self.tau
27 cand=s[:,t].argmax(-1); on=.5+self.delta/2; off=.5-self.delta/2
28 can=(h.gather(1,active[:,None]).squeeze(1)<=off)&(h.gather(1,cand[:,None]).squeeze(1)>=on)&(cand!=active)
29 active=torch.where(can,cand,active); switches+=int(can.sum())
30 choices.append(active); outs.append(h)
31 ch=torch.stack(choices,1); g=F.one_hot(ch,K).float(); g=g+(s-g).detach()
32 self.stats={'switches':switches,'h_deriv':float(torch.stack(outs,1)[:,1:].sub(torch.stack(outs,1)[:,:-1]).abs().max().detach().cpu())}
33 pred=torch.stack([e(z) for e in self.experts],-1).squeeze(-2)
34 return (pred[:,-1,:]*g[:,-1,:]).sum(-1,keepdim=True)
35def seedall(s):
36 random.seed(s); np.random.seed(s); torch.manual_seed(s)
37def run(mode,cfg,seed,collect=False):
38 seedall(seed); ds=get_dataset('sequence',seed,NTR,NTE)
39 m=Router(ds['input_shape'][0],mode,cfg['temperature'],cfg.get('delta',.12),cfg.get('tau',2.))
40 _,metric,_=train_model(m,ds,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *_:None)
41 if collect: return float(metric),dict(m.stats)
42 return float(metric)
43def main():
44 # Union parity: baseline evaluates every lr-temperature pair used by idea.
45 base=sweep_baseline(lambda c:lambda s:run('soft',c,s),GRID,seeds=SEEDS)
46 best=base['best_cfg']; idea_grid=[dict(best,delta=.08,tau=1.5),dict(best,delta=.12,tau=2.),dict(best,delta=.20,tau=3.)]
47 # Baseline's decisive temperature/lr knobs were swept above; idea uses same best shared config.
48 ir=evaluate(lambda s:run('hyst',dict(best,delta=.12,tau=2.),s),SEEDS)
49 sig=[]
50 for s in SEEDS:
51 _,st=run('hyst',dict(best,delta=.12,tau=2.),s,True); sig.append(st)
52 h=np.array([x.get('h_deriv',np.nan) for x in sig]); observed=float(np.nanmax(h))
53 predicted=.12/observed if observed>0 else float('nan')
54 extra={'prediction':'dwell >= delta/L_h','delta_eta':.12,'L_h_observed':observed,'predicted_min_dwell':predicted,'observed_switch_statistics':sig,'confirmed':False,'note':'The bench forward pass records switches but does not expose per-example dwell intervals.'}
55 rep=make_report('sequence','routed_transformer',base,ir,extra)
56 rep['idea']['sweep']=[{'cfg':c,'mean':evaluate(lambda s,c=c:run('hyst',c,s),SEEDS)['mean']} for c in idea_grid]
57 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
58 print(json.dumps(rep,indent=2))
59if __name__=='__main__': main()