Dissipative Softmax Latent Layer / official_bench.py
Failed on benchmark
1import json, random
2import numpy as np
3import torch
4from torch import nn
5import sys
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import get_dataset, make_model, train_model, evaluate, sweep_baseline, make_report
8
9TRACK='token_expert_sequence'
10MODEL='mlp_tiny'
11SEEDS=tuple(range(8))
12EPOCHS=25
13
14class DissipativeNet(nn.Module):
15 """Same mlp_tiny backbone, with a differentiable reset/cycle latent readout."""
16 def __init__(self, input_shape, out_dim=2, reset_rate=2.0, drive=.12, steps=16):
17 super().__init__()
18 # Exact mlp_tiny: 24 -> 64 -> 64 -> 2.
19 dim=int(np.prod(input_shape))
20 self.body=nn.Sequential(nn.Linear(dim,64),nn.ReLU(),nn.Linear(64,64),nn.ReLU(),nn.Linear(64,out_dim))
21 self.r=float(reset_rate); self.drive=float(drive); self.steps=int(steps)
22 def forward(self,x):
23 X=self.body(x.reshape(x.shape[0],-1)); K=X.shape[1]
24 p=torch.softmax(X,dim=1)
25 B=x.shape[0]; rates=torch.zeros(B,K+1,K+1,device=x.device,dtype=x.dtype)
26 rates[:,0,1:]=self.r*p
27 rates[:,1:,0]=self.r
28 d=(X[:,None,:]-X[:,:,None])/2
29 W=torch.exp(d)*(1-torch.eye(K,device=x.device,dtype=x.dtype)[None])
30 rates[:,1:,1:]=W
31 # Directed reset-sector cycle 0 -> 1 -> ... -> K-1 -> 0.
32 for i in range(K-1): rates[:,i,i+1]+=self.drive
33 rates[:,K-1,0]+=self.drive
34 total=rates.sum(2)
35 dt=.45/(total.max().detach()+1e-6)
36 P=torch.eye(K+1,device=x.device,dtype=x.dtype)[None]-dt*torch.diag_embed(total)+dt*rates
37 q=torch.zeros(B,K+1,device=x.device,dtype=x.dtype); q[:,0]=1.
38 for _ in range(self.steps): q=torch.bmm(q[:,None,:],P).squeeze(1).clamp_min(1e-8); q=q/q.sum(1,keepdim=True)
39 return torch.log(q[:,1:].clamp_min(1e-8))
40
41def set_seed(s):
42 random.seed(s); np.random.seed(s); torch.manual_seed(s)
43
44def baseline_fn(cfg):
45 def run(seed):
46 set_seed(seed); d=get_dataset(TRACK,seed,n_train=400,n_test=400)
47 net=nn.Sequential(nn.Flatten(), make_model(MODEL,d['input_shape'],d['out_dim']))
48 _,metric,_=train_model(net,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,log=lambda *_:None)
49 return metric
50 return run
51
52def idea_fn(cfg):
53 def run(seed):
54 set_seed(seed); d=get_dataset(TRACK,seed,n_train=400,n_test=400)
55 net=DissipativeNet(d['input_shape'],d['out_dim'],cfg['r'])
56 _,metric,_=train_model(net,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,log=lambda *_:None)
57 return metric
58 return run
59
60def signature(cfg):
61 vals=[]; currents=[]
62 for seed in SEEDS:
63 set_seed(seed); d=get_dataset(TRACK,seed,n_train=400,n_test=400)
64 net=DissipativeNet(d['input_shape'],d['out_dim'],cfg['r'])
65 net,_,_=train_model(net,d,epochs=cfg['epochs'],lr=cfg['lr'],batch=128,log=lambda *_:None)
66 with torch.no_grad():
67 dev=next(net.parameters()).device; xt=torch.as_tensor(d['xte'],dtype=torch.float32,device=dev)
68 X=net.body(xt.reshape(len(d['xte']),-1)); p=torch.softmax(X,1).cpu().numpy()
69 q=torch.exp(net(xt)).cpu().numpy()
70 vals.append(float(np.abs(q-p).mean()))
71 currents.append(float(np.abs(q[:,0]-q[:,1]).mean()))
72 # A model-behavior check, not an analytical identity: rapid reset should
73 # reduce occupation mismatch relative to the nearby low-reset setting.
74 low=[]
75 for seed in SEEDS:
76 set_seed(seed); d=get_dataset(TRACK,seed,n_train=400,n_test=400)
77 net=DissipativeNet(d['input_shape'],d['out_dim'],.5)
78 net,_,_=train_model(net,d,epochs=EPOCHS,lr=cfg['lr'],batch=128,log=lambda *_:None)
79 with torch.no_grad():
80 dev=next(net.parameters()).device; xt=torch.as_tensor(d['xte'],dtype=torch.float32,device=dev)
81 X=net.body(xt.reshape(len(d['xte']),-1)); p=torch.softmax(X,1); q=torch.exp(net(xt))
82 low.append(float(torch.abs(q-p).mean()))
83 return {'predicted_effect':'occupation mismatch decreases as reset_rate increases', 'mismatch_at_r':float(np.mean(vals)), 'mismatch_at_r_0.5':float(np.mean(low)), 'observed_cycle_proxy':float(np.mean(currents)), 'confirmed':bool(np.mean(vals)<np.mean(low))}
84
85def main():
86 # Union parity: all idea learning rates are also baseline candidates.
87 grid=[{'lr':lr,'epochs':EPOCHS} for lr in (1e-3,3e-3,1e-2)]
88 base=sweep_baseline(baseline_fn,grid,seeds=(0,1,2,3))
89 # Three idea settings at the selected baseline lr; reset is the method knob.
90 lr=base['best_cfg']['lr']
91 idea_cfgs=[{'lr':lr,'epochs':EPOCHS,'r':r} for r in (.5,2.,8.)]
92 idea_runs=[(cfg,evaluate(idea_fn(cfg),seeds=SEEDS)) for cfg in idea_cfgs]
93 best_cfg,idea_res=min(idea_runs,key=lambda z:z[1]['mean'])
94 extra=signature(best_cfg)
95 rep=make_report(TRACK,MODEL,base,idea_res,extra=extra)
96 rep['idea_sweep']=[{'cfg':cfg,'result':res} for cfg,res in idea_runs]
97 rep['custom_track']={'name':TRACK,'file':'/home/maxwelhelp/all/math2nn/bench/custom_tracks/token_expert_sequence.py','domain':'moe-routing'}
98 with open('bench_report.json','w') as f: json.dump(rep,f,indent=2)
99 print(json.dumps(rep,indent=2))
100if __name__=='__main__': main()