Van der Corput progressive expert scheduler / vdc_bench.py
Mechanism confirmed, baseline not beaten
1import sys, math, json
2from pathlib import Path
3import numpy as np
4import torch
5import torch.nn as nn
6
7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
8from bench import make_report, permutation_pvalue
9
10META = {'name': 'expert_balanced_regression', 'domain': 'moe-routing', 'description': 'Regression with a top-1 mixture of equally sized experts; routing balance is the intervention.'}
11
12
13def get_dataset(seed, n_train=400, n_test=400):
14 rng = np.random.default_rng(seed)
15 n = n_train + n_test
16 x = rng.uniform(-1, 1, (n, 8)).astype(np.float32)
17 y = (np.sin(3*x[:, 0]) + .45*x[:, 1]**2 - .35*x[:, 2]*x[:, 3] + .2*x[:, 4] - .15*x[:, 5]**3).astype(np.float32)
18 y += rng.normal(0, .025, n).astype(np.float32)
19 return {'xtr': x[:n_train], 'ytr': y[:n_train,None], 'xte': x[n_train:], 'yte': y[n_train:,None], 'task':'regression', 'metric':'mse', 'out_dim':1}
20
21
22def vdc(m):
23 z, f = 0.0, 0.5
24 while m:
25 z += (m & 1) * f; m >>= 1; f *= .5
26 return z
27
28
29def routes(mode, start, n, D, logits=None, rho=1.0):
30 rng = np.random.default_rng(start + 991)
31 out=[]
32 for i in range(n):
33 if mode == 'vdc' and (i == 0 or rng.random() < rho): out.append(int(D*vdc(start+i)) % D)
34 elif mode == 'random': out.append(int(rng.integers(D)))
35 else: out.append(int(torch.argmax(logits[i]).item()) if logits is not None else int(rng.integers(D)))
36 return np.asarray(out, dtype=np.int64)
37
38
39class Top1MoE(nn.Module):
40 def __init__(self, mode='vdc', D=8, rho=1.0):
41 super().__init__(); self.mode, self.D, self.rho = mode, D, rho
42 self.router = nn.Linear(8, D)
43 self.experts = nn.ModuleList([nn.Sequential(nn.Linear(8,32),nn.ReLU(),nn.Linear(32,1)) for _ in range(D)])
44 self.assignment_index = 0
45 self.last_routes = None
46 def forward(self, x):
47 logits = self.router(x)
48 n = x.shape[0]
49 if self.mode == 'vdc':
50 r = routes('vdc', self.assignment_index, n, self.D, rho=self.rho)
51 elif self.mode == 'random':
52 r = routes('random', self.assignment_index, n, self.D)
53 else: r = logits.argmax(1).detach().cpu().numpy()
54 self.assignment_index += n
55 self.last_routes = r
56 y = torch.empty((n,1), device=x.device)
57 for e in range(self.D):
58 ix = np.flatnonzero(r == e)
59 if len(ix): y[ix] = self.experts[e](x[ix])
60 return y
61
62
63def train(seed, mode, lr, epochs, rho=1.0, return_sig=False):
64 torch.manual_seed(seed); np.random.seed(seed)
65 d=get_dataset(seed); device='cuda' if torch.cuda.is_available() else 'cpu'
66 model=Top1MoE(mode, rho=rho).to(device)
67 x=torch.tensor(d['xtr'],device=device); y=torch.tensor(d['ytr'],device=device)
68 opt=torch.optim.Adam(model.parameters(),lr=lr)
69 for _ in range(epochs):
70 model.train(); perm=torch.randperm(len(x),device=device)
71 for j in range(0,len(x),128):
72 ix=perm[j:j+128]; loss=((model(x[ix])-y[ix])**2).mean()
73 opt.zero_grad(); loss.backward(); opt.step()
74 model.eval(); xt=torch.tensor(d['xte'],device=device); yt=torch.tensor(d['yte'],device=device)
75 with torch.no_grad(): metric=float(((model(xt)-yt)**2).mean())
76 # Re-test trained model behavior: assignment prefix discrepancy on a fixed evaluation stream.
77 model.assignment_index=0; counts=np.zeros(model.D,int); max_disc=0.0; total=0
78 with torch.no_grad():
79 for j in range(0,len(xt),128):
80 model(xt[j:j+128]); rr=model.last_routes
81 for e in rr:
82 counts[e]+=1; total+=1
83 max_disc=max(max_disc,float(np.max(np.abs(counts-total/model.D))))
84 return metric, {'max_leaf_prefix_discrepancy':max_disc,'final_load_variance':float(np.var(counts))} if return_sig else None
85
86
87def eval_cfg(mode, cfg, seeds=range(8), sig=False):
88 vals=[]; signatures=[]
89 for s in seeds:
90 m,q=train(s,mode,cfg['lr'],cfg['epochs'],cfg.get('rho',1.0),sig); vals.append(m)
91 if q: signatures.append(q)
92 return {'per_seed':[float(v) for v in vals], 'mean':float(np.mean(vals)), 'std':float(np.std(vals,ddof=1)), 'config':cfg, 'signatures':signatures}
93
94
95def main():
96 # Union grid parity: all idea learning rates are included in baseline sweep.
97 grid=[{'lr':lr,'epochs':12,'rho':rho} for lr in (1e-3,3e-3,1e-2) for rho in (0.5,1.0)]
98 sweep=[]
99 for cfg in grid:
100 r=eval_cfg('random',cfg,seeds=range(4)); sweep.append({'config':cfg,'mean':r['mean'],'per_seed':r['per_seed']})
101 best=min(sweep,key=lambda z:z['mean'])['config']
102 base=eval_cfg('random',best,seeds=range(8))
103 idea_runs=[eval_cfg('vdc',cfg,seeds=range(8),sig=True) for cfg in grid]
104 idea=min(idea_runs,key=lambda z:z['mean'])
105 diffs=[a-b for a,b in zip(idea['per_seed'],base['per_seed'])]
106 cmp={'delta_mean':float(np.mean(diffs)),'idea_wins':sum(x<0 for x in diffs),'n_pairs':8,'per_seed_diffs':diffs,'p_value':permutation_pvalue(diffs)}
107 if cmp['delta_mean']<0 and cmp['p_value']<.05: cmp['verdict']='idea better (significant)'
108 elif cmp['delta_mean']>0 and cmp['p_value']<.05: cmp['verdict']='idea worse (significant)'
109 else: cmp['verdict']='no significant win'
110 cmp['system_worked']=cmp['verdict']=='idea better (significant)'
111 sig=idea['signatures']; obs=float(np.mean([q['max_leaf_prefix_discrepancy'] for q in sig])); pred=1.0
112 report={'bench_version':1,'track':'expert_balanced_regression','model':'top1_moe_shared','metric_direction':'lower is better','n_seeds':8,'baseline':{'best_cfg':best,'sweep':sweep,'full':base},'idea':idea,'comparison':cmp,'mechanism_signature':{'prediction':'VDC trained-model assignment prefixes have leaf discrepancy <=1, unlike random routing','predicted_max_discrepancy':pred,'observed_mean_max_discrepancy':obs,'confirmed':bool(obs<=1.05)},'custom_track':{'name':'expert_balanced_regression','file':'vdc_bench.py','domain':'moe-routing'}}
113 Path('bench_report.json').write_text(json.dumps(report,indent=2)); print(json.dumps(report,indent=2))
114
115if __name__=='__main__': main()