Expansion-balanced MoE routing / bench_expansion.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 1import sys, math, json, random
 2from pathlib import Path
 3import numpy as np
 4import torch
 5import torch.nn as nn
 6import torch.nn.functional as F
 7sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
 8from bench import get_dataset, sweep_baseline, make_report
 9
10SEEDS=tuple(range(8)); E=8; K=4; EPS=1.0
11
12class RoutedForecaster(nn.Module):
13    def __init__(self, win, use_expand, lam, temperature):
14        super().__init__(); self.use_expand=use_expand; self.lam=lam
15        self.inp=nn.Linear(1,16); self.pos=nn.Parameter(torch.zeros(1,win,16)); nn.init.normal_(self.pos,.02)
16        self.enc=nn.TransformerEncoder(nn.TransformerEncoderLayer(16,2,32,batch_first=True,dropout=0.),1)
17        self.router=nn.Linear(16,E)
18        self.experts=nn.ModuleList([nn.Sequential(nn.Linear(16,16),nn.ReLU(),nn.Linear(16,16)) for _ in range(E)])
19        self.head=nn.Linear(16,1); self.temperature=temperature; self.last_sig={}
20    def expansion_loss(self,p):
21        n=p.shape[0]; rng=np.random.default_rng(1771); vals=[]; hard=[]; bounds=[]
22        for m in [K,2*K,4*K,8*K]:
23            if m<=n//2:
24                for _ in range(1):
25                    ix=torch.as_tensor(rng.choice(n,m,False),device=p.device); c=1-torch.prod(1-p[ix],0); C=c.sum()
26                    b=EPS*m/(math.log(3*m/K)**2); vals.append(F.relu(torch.as_tensor(b,device=p.device)-C)**2)
27                    hard.append(int(p[ix].argmax(1).unique().numel())); bounds.append(b)
28        return torch.stack(vals).mean(),hard,bounds
29    def forward(self,x):
30        h=self.enc(self.inp(x.unsqueeze(-1))+self.pos[:,:x.shape[1]])
31        p=F.softmax(self.router(h)/self.temperature,-1); ys=torch.stack([e(h) for e in self.experts],2)
32        out=self.head((ys*p.unsqueeze(-1)).sum(2)[:,-1]); load=p.mean((0,1)); lb=E*(load**2).sum()
33        ex,hard,bounds=self.expansion_loss(p.reshape(-1,E))
34        counts=p.argmax(-1).reshape(-1).bincount(minlength=E).float()
35        self.last_sig={'soft_C':float((1-torch.prod(1-p.reshape(-1,E),0)).sum().detach()),'mean_hard_neighborhood':float(np.mean(hard)),'violation_fraction':float(np.mean([a<b for a,b in zip(hard,bounds)])),'load_std':float(counts.std().detach())}
36        return out,lb,ex
37
38def train_one(seed,cfg,use_expand,return_sig=False):
39    random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
40    d=get_dataset('sequence',seed,n_train=200,n_test=100)
41    # Small CPU path is deterministic and avoids shared-GPU contention; this is the harness fallback path.
42    dev=torch.device('cpu')
43    net=RoutedForecaster(d['input_shape'][0],use_expand,cfg['lambda'],cfg['temperature']).to(dev)
44    opt=torch.optim.Adam(net.parameters(),lr=cfg['lr']); x=d['xtr'].to(dev); y=d['ytr'].reshape(-1,1).to(dev)
45    for ep in range(cfg['epochs']):
46        for ix in torch.randperm(len(x)).split(128):
47            pred,lb,ex=net(x[ix]); warm=min(1.,(ep+1)/max(1,cfg['epochs']//10))
48            loss=F.mse_loss(pred,y[ix])+.01*(lb-1)**2+(cfg['lambda']*warm*ex if use_expand else 0.)
49            opt.zero_grad(); loss.backward(); opt.step()
50    with torch.no_grad(): pred,_,_=net(d['xte'].to(dev)); metric=float(F.mse_loss(pred,d['yte'].reshape(-1,1).to(dev)))
51    return (metric,net.last_sig) if return_sig else metric
52
53def main():
54    base_grid=[{'lr':lr,'temperature':t,'lambda':0.,'epochs':3} for lr in [.001,.003] for t in [.7,1.0]]
55    base=sweep_baseline(lambda c: lambda s: train_one(s,c,False),base_grid,seeds=(0,1,2,3))
56    best=base['best_cfg']; idea_grid=[dict(best, **{'lambda': v}) for v in [.01,.03,.08]]
57    # Three-config idea sweep on the same lr/temp selected by the baseline sweep.
58    idea_per=[]; idea_sigs=[]
59    for s in SEEDS:
60        rr=[train_one(s,c,True,True) for c in idea_grid]; j=int(np.argmin([q[0] for q in rr])); idea_per.append(rr[j][0]); idea_sigs.append(rr[j][1])
61    idea={'mean':float(np.mean(idea_per)),'std':float(np.std(idea_per)),'per_seed':idea_per,'n':8,'sweep':[{'cfg':c,'mean':float(np.mean([train_one(s,c,True) for s in (0,1,2,3)]))} for c in idea_grid]}
62    base_sigs=[train_one(s,best,False,True)[1] for s in SEEDS]
63    sig={'prediction':'expansion penalty increases token-group expert neighborhood and reduces load dispersion','baseline_observed':{k:float(np.mean([a[k] for a in base_sigs])) for k in base_sigs[0]},'idea_observed':{k:float(np.mean([a[k] for a in idea_sigs])) for k in idea_sigs[0]},'confirmed':bool(np.mean([a['mean_hard_neighborhood'] for a in idea_sigs])>np.mean([a['mean_hard_neighborhood'] for a in base_sigs]) and np.mean([a['load_std'] for a in idea_sigs])<np.mean([a['load_std'] for a in base_sigs]))}
64    rep=make_report('sequence','transformer_tiny',base,idea,sig); rep['idea_hyperparameter_grid']=idea_grid; rep['notes']='Matched compact sequence forecaster; same 8-expert dense MoE and training budget; CPU fallback used to avoid shared GPU contention.'
65    Path('bench_report.json').write_text(json.dumps(rep,indent=2)); print(json.dumps(rep,indent=2))
66if __name__=='__main__': main()