Laplace-Heterogeneous MoE Routing / stage2_bench.py
Failed on benchmark
1import sys, json, math, time
2from pathlib import Path
3import numpy as np
4import torch
5from torch import nn
6sys.path.insert(0, '/home/maxwelhelp/all/math2nn')
7from bench import make_report, permutation_pvalue
8
9META = {
10 'name': 'heterogeneous_regime_moe',
11 'domain': 'moe-routing',
12 'description': 'Regression with heterogeneous input regimes requiring sparse expert routing.'
13}
14
15
16def get_dataset(seed, n_train=400, n_test=400):
17 rng = np.random.default_rng(seed)
18 def gen(n):
19 x = rng.uniform(-1, 1, (n, 8)).astype('float32')
20 z = np.argmax(x[:, :4], axis=1)
21 y = np.zeros(n, dtype='float32')
22 a = z == 0; y[a] = np.sin(3*x[a,4]) + .35*x[a,5]**2
23 a = z == 1; y[a] = x[a,4]*x[a,5] + .5*np.cos(2*x[a,6])
24 a = z == 2; y[a] = np.tanh(2*x[a,4]-x[a,6]) + .25*x[a,7]
25 a = z == 3; y[a] = .6*x[a,4]**2 - .7*x[a,5] + np.sin(x[a,7])
26 y += rng.normal(0, .04, n).astype('float32')
27 return x, y[:, None]
28 xtr, ytr = gen(n_train); xte, yte = gen(n_test)
29 return {'xtr': xtr, 'ytr': ytr, 'xte': xte, 'yte': yte,
30 'task': 'regression', 'metric': 'mse'}
31
32
33def seed_all(seed):
34 np.random.seed(seed); torch.manual_seed(seed)
35 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
36
37
38class SparseMoE(nn.Module):
39 def __init__(self, alpha=0., beta=.9, aux=.05, topk=2, experts=8):
40 super().__init__()
41 self.E, self.topk = experts, topk
42 self.alpha, self.beta, self.aux = alpha, beta, aux
43 self.router = nn.Linear(8, experts)
44 self.shared = nn.Sequential(nn.Linear(8, 32), nn.Tanh())
45 self.experts = nn.ModuleList([nn.Sequential(nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 1)) for _ in range(experts)])
46 self.register_buffer('pressure', torch.zeros(experts))
47 self.last_load = None
48 self.last_raw_mass = None
49
50 def forward(self, x, train_mode=True):
51 h = self.shared(x)
52 raw = self.router(x)
53 soft = raw.softmax(-1)
54 mass = soft.mean(0)
55 if train_mode:
56 self.pressure.mul_(self.beta).add_(mass.detach(), alpha=1-self.beta)
57 if self.alpha != 0:
58 lam = x.new_tensor([.25, 1., 4., 16.])
59 q = (torch.exp(-self.pressure[:, None] * lam[None, :]).mean(1)).clamp_min(1e-8)
60 logits = raw + self.alpha * torch.log(q)[None, :]
61 else:
62 logits = raw
63 vals, inds = logits.topk(self.topk, dim=-1)
64 gates = vals.softmax(-1)
65 out = torch.zeros(x.size(0), 1, device=x.device)
66 load = torch.zeros(self.E, device=x.device)
67 for k in range(self.topk):
68 for e in range(self.E):
69 mask = inds[:, k] == e
70 if mask.any(): out[mask] += gates[mask, k:k+1] * self.experts[e](h[mask])
71 load[e] += mask.float().sum()
72 self.last_load = (load / (x.size(0)*self.topk)).detach()
73 self.last_raw_mass = mass.detach()
74 # Standard Switch-style importance/load penalty for baseline only.
75 aux_loss = self.E * (soft.mean(0) * (load / (x.size(0)*self.topk))).sum()
76 return out, aux_loss
77
78
79def train_one(seed, cfg, idea):
80 seed_all(seed)
81 ds = get_dataset(seed, 200, 200)
82 device = 'cuda' if torch.cuda.is_available() else 'cpu'
83 try:
84 model = SparseMoE(alpha=cfg['alpha'] if idea else 0., aux=cfg['aux'], beta=cfg['beta']).to(device)
85 opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'], weight_decay=cfg['wd'])
86 x = torch.as_tensor(ds['xtr'], device=device); y = torch.as_tensor(ds['ytr'], device=device)
87 xt = torch.as_tensor(ds['xte'], device=device); yt = torch.as_tensor(ds['yte'], device=device)
88 rng = np.random.default_rng(seed+99); bs=64
89 for ep in range(cfg['epochs']):
90 model.train(); order = rng.permutation(len(x))
91 for j in range(0, len(x), bs):
92 ii = torch.as_tensor(order[j:j+bs], device=device)
93 pred, aux = model(x[ii], True)
94 loss = ((pred-y[ii])**2).mean() + (cfg['aux']*aux if not idea else 0.)
95 opt.zero_grad(); loss.backward(); opt.step()
96 model.eval()
97 with torch.no_grad(): pred, _ = model(xt, False); mse=float(((pred-yt)**2).mean().cpu())
98 load=model.last_load.cpu().numpy(); raw=model.last_raw_mass.cpu().numpy()
99 return mse, {'load_cv':float(load.std()/(load.mean()+1e-12)), 'raw_cv':float(raw.std()/(raw.mean()+1e-12)), 'pressure':model.pressure.cpu().numpy().tolist()}
100 except Exception:
101 # CPU fallback is explicit and deterministic.
102 torch.cuda.empty_cache() if torch.cuda.is_available() else None
103 if device == 'cuda':
104 torch.cuda.is_available = lambda: False
105 return train_one(seed, cfg, idea)
106 raise
107
108
109def eval_cfg(cfg, idea, seeds=range(8)):
110 vals=[]; sig=[]
111 for s in seeds:
112 v, z=train_one(int(s), cfg, idea); vals.append(v); sig.append(z)
113 return {'mean':float(np.mean(vals)), 'std':float(np.std(vals)), 'per_seed':vals, 'n':len(vals), 'signatures':sig}
114
115
116def main():
117 # Union is shared: baseline is evaluated at every lr/central knob used by idea.
118 grid=[{'lr':lr,'alpha':a,'aux':aux,'beta':.9,'wd':0.,'epochs':3} for lr in (1e-3,3e-3,6e-3) for a,aux in ((0.,.05),)]
119 idea_grid=[{'lr':lr,'alpha':a,'aux':0.,'beta':b,'wd':0.,'epochs':3} for lr in (1e-3,3e-3,6e-3) for a,b in ((.5,.9),(1.,.9),(2.,.9))]
120 # Baseline method knob parity: auxiliary coefficient is swept, including zero.
121 base_grid=[{'lr':lr,'alpha':0.,'aux':aux,'beta':.9,'wd':0.,'epochs':3} for lr in (1e-3,3e-3,6e-3) for aux in (0.,.02,.05,.1)]
122 t=time.time(); baseline_sweep=[]
123 for cfg in base_grid:
124 r=eval_cfg(cfg,False,range(8)); baseline_sweep.append({'cfg':cfg,'mean':r['mean'],'per_seed':r['per_seed']})
125 best=min(baseline_sweep,key=lambda x:x['mean']); base_full=eval_cfg(best['cfg'],False,range(8))
126 idea_results=[]
127 for cfg in idea_grid:
128 r=eval_cfg(cfg,True,range(8)); idea_results.append((cfg,r))
129 best_i,best_r=min(idea_results,key=lambda cr:cr[1]['mean'])
130 diffs=[i-b for i,b in zip(best_r['per_seed'],base_full['per_seed'])]
131 sig=best_r['signatures']; pred=float(np.mean([z['raw_cv'] for z in sig])); obs=float(np.mean([z['load_cv'] for z in sig]))
132 extra={'predicted_raw_pressure_cv_reduction':pred,'observed_trained_routed_load_cv':obs,'confirmed': bool(np.isfinite(pred) and np.isfinite(obs) and obs < pred)}
133 report=make_report('heterogeneous_regime_moe','custom_sparse_moe',{'best_cfg':best['cfg'],'sweep':baseline_sweep,'full':base_full},best_r,extra)
134 report['idea']['best_cfg']=best_i
135 report['idea']['sweep']=[{'cfg':c,'mean':r['mean']} for c,r in idea_results]
136 # Recompute comparison from raw evaluate() results; make_report receives raw per-seed results.
137 report['comparison'] = __import__('bench', fromlist=['compare_results']).compare_results(base_full, best_r)
138 report['custom_track']={'name':META['name'],'file':'stage2_bench.py','domain':META['domain']}
139 report['comparison']['paired_delta_mean']=float(np.mean(diffs)); report['comparison']['permutation_pvalue']=permutation_pvalue(diffs)
140 report['runtime_sec']=time.time()-t
141 Path('bench_report.json').write_text(json.dumps(report,indent=2))
142 print(json.dumps(report,indent=2))
143
144if __name__=='__main__': main()