Van der Corput progressive expert scheduler / vdc_bench.py

Mechanism confirmed, baseline not beaten

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