Van der Corput progressive expert scheduler / registered_bench.py

Mechanism confirmed, baseline not beaten

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