Van der Corput progressive expert scheduler / registered_bench.py
Mechanism confirmed, baseline not beaten
1import sys, json
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, sweep_baseline, evaluate, make_report
8
9TRACK = 'expert_balanced_regression'
10D = 8
11
12def vdc(m):
13 z, f = 0.0, 0.5
14 while m:
15 z += (m & 1) * f
16 m >>= 1; f *= 0.5
17 return z
18
19def route(mode, start, n, D, x=None, rho=1.0):
20 rng = np.random.default_rng(start + 17015)
21 if mode == 'vdc':
22 return np.asarray([int(D*vdc(start+i)) % D if (i == 0 or rng.random() < rho) else int(rng.integers(D)) for i in range(n)])
23 if mode == 'random':
24 return rng.integers(0, D, size=n)
25 return x.argmax(1).detach().cpu().numpy()
26
27class MoE(nn.Module):
28 def __init__(self, mode, rho):
29 super().__init__(); self.mode=mode; self.rho=rho; self.pos=0
30 self.router=nn.Linear(8,D)
31 self.experts=nn.ModuleList([nn.Sequential(nn.Linear(8,32),nn.ReLU(),nn.Linear(32,1)) for _ in range(D)])
32 self.last=None
33 def forward(self,x):
34 logits=self.router(x); n=len(x)
35 if self.mode == 'greedy': r=route('greedy',self.pos,n,D,logits)
36 else: r=route(self.mode,self.pos,n,D,rho=self.rho)
37 self.pos += n; self.last=r
38 y=torch.empty((n,1),device=x.device)
39 for e in range(D):
40 ii=np.flatnonzero(r==e)
41 if len(ii): y[ii]=self.experts[e](x[ii])
42 return y
43
44def run(seed, mode, cfg, signature=False):
45 torch.manual_seed(seed); np.random.seed(seed)
46 d=get_dataset(TRACK, seed=seed, n_train=400, n_test=400)
47 dev='cuda' if torch.cuda.is_available() else 'cpu'
48 try:
49 net=MoE(mode,cfg.get('rho',1.0)).to(dev)
50 x,y=d['xtr'].to(dev),d['ytr'].to(dev)
51 opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'])
52 for _ in range(cfg['epochs']):
53 p=torch.randperm(len(x),device=dev)
54 for j in range(0,len(x),128):
55 ii=p[j:j+128]; loss=((net(x[ii])-y[ii])**2).mean()
56 opt.zero_grad(); loss.backward(); opt.step()
57 net.eval(); xt,yt=d['xte'].to(dev),d['yte'].to(dev)
58 with torch.no_grad(): metric=float(((net(xt)-yt)**2).mean())
59 if not signature: return metric
60 net.pos=0; counts=np.zeros(D,int); mx=0.; total=0
61 with torch.no_grad():
62 for j in range(0,len(xt),128):
63 net(xt[j:j+128])
64 for e in net.last:
65 counts[e]+=1; total+=1
66 mx=max(mx,float(np.max(np.abs(counts-total/D))))
67 return metric, {'max_leaf_prefix_discrepancy':mx,'final_load_variance':float(np.var(counts))}
68 except RuntimeError:
69 dev='cpu'; net=MoE(mode,cfg.get('rho',1.0)).to(dev)
70 x,y=d['xtr'],d['ytr']; opt=torch.optim.Adam(net.parameters(),lr=cfg['lr'])
71 for _ in range(cfg['epochs']):
72 p=torch.randperm(len(x))
73 for j in range(0,len(x),128):
74 ii=p[j:j+128]; loss=((net(x[ii])-y[ii])**2).mean(); opt.zero_grad(); loss.backward(); opt.step()
75 net.eval()
76 with torch.no_grad(): return float(((net(d['xte'])-d['yte'])**2).mean())
77
78def fn(mode,cfg,signature=False):
79 return lambda seed: run(seed,mode,cfg,signature)
80
81def main():
82 # Every idea learning rate is present in the baseline grid (union parity).
83 grid=[{'lr':lr,'epochs':12,'rho':rho} for lr in (0.001,0.003,0.01) for rho in (0.5,1.0)]
84 base=sweep_baseline(lambda c: fn('random',c),grid)
85 # Three idea settings: baseline best and two nearby method settings.
86 idea_cfgs=[base['best_cfg'],{'lr':0.003,'epochs':12,'rho':1.0},{'lr':0.01,'epochs':12,'rho':1.0}]
87 ideas=[evaluate(fn('vdc',c,False)) for c in idea_cfgs]
88 best=min(zip(idea_cfgs,ideas),key=lambda z:z[1]['mean'])
89 cfg,idea=best; idea['config']=cfg
90 sigs=[run(s,'vdc',cfg,True)[1] for s in range(8)]
91 extra={'prediction':'trained VDC assignments maintain max leaf prefix discrepancy <=1','predicted_max_discrepancy':1.0,'observed_mean_max_discrepancy':float(np.mean([s['max_leaf_prefix_discrepancy'] for s in sigs])),'observed_final_load_variance':float(np.mean([s['final_load_variance'] for s in sigs])),'confirmed':bool(max(s['max_leaf_prefix_discrepancy'] for s in sigs)<=1.05)}
92 report=make_report(TRACK,'top1_moe_shared',base,idea,extra)
93 report['custom_track']={'name':TRACK,'file':'expert_balanced_regression.py','domain':'moe-routing'}
94 Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
95if __name__=='__main__': main()