Expansion-balanced MoE routing / bench_expansion.py
Mechanism confirmed, baseline not beaten
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()