import sys,json,random import numpy as np, torch import torch.nn as nn import torch.nn.functional as F sys.path.insert(0,'/home/maxwelhelp/all/math2nn') from bench import get_dataset,train_model,evaluate,sweep_baseline,make_report SEEDS=tuple(range(8)); EPOCHS=10; NTR=400; NTE=200; K=4 GRID=[{'lr':.0015,'temperature':.7},{'lr':.003,'temperature':1.0},{'lr':.006,'temperature':1.4}] class Router(nn.Module): def __init__(self,win,mode,temp=1.,delta=.12,tau=2.): super().__init__(); self.mode=mode; self.temp=temp; self.delta=delta; self.tau=tau d=48; self.inp=nn.Linear(1,d); self.pos=nn.Parameter(torch.randn(1,win,d)*.02) self.enc=nn.TransformerEncoder(nn.TransformerEncoderLayer(d,2,96,batch_first=True,dropout=0),2) self.score=nn.Linear(d,K); self.experts=nn.ModuleList([nn.Linear(d,1) for _ in range(K)]) self.stats={} def forward(self,x): z=self.enc(self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]]) s=F.softmax(self.score(z)/self.temp,-1) if self.mode=='soft': g=s; ch=g.argmax(-1) else: # Differentiable leaky hysteresis approximation: cumulative state is detached # for routing decisions, while the soft scores provide the straight-through gradient. b=x.shape[0]; h=torch.full((b,K),1./K,device=x.device); active=torch.zeros(b,dtype=torch.long,device=x.device) outs=[]; choices=[]; switches=0; min_dwell=10**9; last=torch.zeros(b,dtype=torch.long,device=x.device) for t in range(x.shape[1]): h=h+(s[:,t]-h)/self.tau cand=s[:,t].argmax(-1); on=.5+self.delta/2; off=.5-self.delta/2 can=(h.gather(1,active[:,None]).squeeze(1)<=off)&(h.gather(1,cand[:,None]).squeeze(1)>=on)&(cand!=active) active=torch.where(can,cand,active); switches+=int(can.sum()) choices.append(active); outs.append(h) ch=torch.stack(choices,1); g=F.one_hot(ch,K).float(); g=g+(s-g).detach() self.stats={'switches':switches,'h_deriv':float(torch.stack(outs,1)[:,1:].sub(torch.stack(outs,1)[:,:-1]).abs().max().detach().cpu())} pred=torch.stack([e(z) for e in self.experts],-1).squeeze(-2) return (pred[:,-1,:]*g[:,-1,:]).sum(-1,keepdim=True) def seedall(s): random.seed(s); np.random.seed(s); torch.manual_seed(s) def run(mode,cfg,seed,collect=False): seedall(seed); ds=get_dataset('sequence',seed,NTR,NTE) m=Router(ds['input_shape'][0],mode,cfg['temperature'],cfg.get('delta',.12),cfg.get('tau',2.)) _,metric,_=train_model(m,ds,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *_:None) if collect: return float(metric),dict(m.stats) return float(metric) def main(): # Union parity: baseline evaluates every lr-temperature pair used by idea. base=sweep_baseline(lambda c:lambda s:run('soft',c,s),GRID,seeds=SEEDS) 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.)] # Baseline's decisive temperature/lr knobs were swept above; idea uses same best shared config. ir=evaluate(lambda s:run('hyst',dict(best,delta=.12,tau=2.),s),SEEDS) sig=[] for s in SEEDS: _,st=run('hyst',dict(best,delta=.12,tau=2.),s,True); sig.append(st) h=np.array([x.get('h_deriv',np.nan) for x in sig]); observed=float(np.nanmax(h)) predicted=.12/observed if observed>0 else float('nan') 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.'} rep=make_report('sequence','routed_transformer',base,ir,extra) rep['idea']['sweep']=[{'cfg':c,'mean':evaluate(lambda s,c=c:run('hyst',c,s),SEEDS)['mean']} for c in idea_grid] with open('bench_report.json','w') as f: json.dump(rep,f,indent=2) print(json.dumps(rep,indent=2)) if __name__=='__main__': main()