Hysteretic Multiscale Sequence Router / bench_hysteretic.py

Failed on benchmark

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