Shell-Wise Balanced MoE Routing / stage2_shell_moe.py
Mechanism confirmed, baseline not beaten
1import sys, json, random
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, get_dataset as bench_get_dataset
9from shell_moe_bench_track import META
10
11ROOT = Path(__file__).parent
12SEEDS = tuple(range(8))
13SWEEP_SEEDS = tuple(range(4))
14E, TOK, D = 8, 16, 4
15
16def seed_all(seed):
17 random.seed(seed); np.random.seed(seed); torch.manual_seed(seed)
18 if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed)
19
20def device():
21 return torch.device('cuda' if torch.cuda.is_available() else 'cpu')
22
23class MoE(nn.Module):
24 def __init__(self, hidden=32):
25 super().__init__()
26 self.router = nn.Sequential(nn.Linear(D, hidden), nn.Tanh(), nn.Linear(hidden, E))
27 self.experts = nn.ModuleList([nn.Sequential(nn.Linear(D, hidden), nn.ReLU(), nn.Linear(hidden, 1)) for _ in range(E)])
28 self.last_stats = {}
29 def forward(self, x, balanced=False, shells=4):
30 # x: batch, tokens, features. Return mean token output and routing logits/mask.
31 z = self.router(x)
32 pref = z.argmax(-1)
33 if balanced:
34 route = shell_route(z.detach().cpu().numpy(), shells)
35 route = torch.as_tensor(route, device=x.device, dtype=torch.long)
36 else:
37 route = pref
38 outs = torch.stack([m(x)[:, :, 0] for m in self.experts], dim=-1)
39 chosen = outs.gather(-1, route.unsqueeze(-1)).squeeze(-1)
40 self.last_stats = routing_stats(z.detach(), route.detach(), pref.detach(), shells)
41 return chosen.mean(1, keepdim=True), z, route
42
43def shell_route(logits, K):
44 b, t, e = logits.shape
45 flat = logits.reshape(-1, e)
46 conf = flat.max(1)
47 order = np.argsort(conf, kind='stable')
48 sh = np.empty(len(conf), dtype=np.int64)
49 for j, i in enumerate(order): sh[i] = min(K-1, j*K//len(conf))
50 route = np.empty(len(conf), dtype=np.int64)
51 for k in range(K):
52 ids = np.flatnonzero(sh == k); n = len(ids)
53 if not n: continue
54 target = np.full(e, n//e, dtype=int)
55 # distribute remainders to experts with greatest preferred demand
56 pref = flat[ids].argmax(1)
57 demand = np.bincount(pref, minlength=e)
58 for q in np.argsort(-demand, kind='stable')[:n % e]: target[q] += 1
59 # greedy maximum-weight quota assignment; each token gets one expert
60 pairs = sorted(((float(flat[i, q]), i, q) for i in ids for q in range(e)), reverse=True)
61 used = np.zeros(e, dtype=int); assigned = set()
62 for val, i, q in pairs:
63 if i not in assigned and used[q] < target[q]:
64 route[i] = q; used[q] += 1; assigned.add(i)
65 assert len(assigned) == n
66 return route.reshape(b, t)
67
68def routing_stats(z, route, pref, K):
69 a = z.reshape(-1, E).cpu().numpy(); r = route.reshape(-1).cpu().numpy(); p = pref.reshape(-1).cpu().numpy()
70 conf = a.max(1); order = np.argsort(conf, kind='stable'); sh = np.empty(len(conf), int)
71 for j, i in enumerate(order): sh[i] = min(K-1, j*K//len(conf))
72 loads = np.bincount(r, minlength=E).astype(float); cv = float(loads.std()/loads.mean())
73 spreads = []
74 for k in range(K):
75 q = np.bincount(r[sh == k], minlength=E)
76 if len(q): spreads.append(int(q.max()-q.min()))
77 return {'load_cv': cv, 'loads': loads.tolist(), 'max_shell_count_spread': max(spreads or [0]),
78 'reassigned_fraction': float(np.mean(r != p)), 'shells': K}
79
80def train_one(ds, seed, lr, aux, balanced, K, epochs=12):
81 seed_all(seed); dev = device()
82 net = MoE().to(dev); opt = torch.optim.Adam(net.parameters(), lr=lr)
83 x = torch.as_tensor(ds['xtr'], dtype=torch.float32, device=dev); y = torch.as_tensor(ds['ytr'], dtype=torch.float32, device=dev)
84 for ep in range(epochs):
85 net.train(); perm = torch.randperm(len(x), device=dev)
86 for ii in range(0, len(x), 128):
87 ix = perm[ii:ii+128]; pred, z, route = net(x[ix], balanced=balanced, shells=K)
88 loss = ((pred-y[ix])**2).mean()
89 if not balanced and aux > 0:
90 prob = z.softmax(-1); hard = torch.nn.functional.one_hot(route, E).float()
91 # standard Switch-style importance/load auxiliary loss
92 loss = loss + aux * E * (prob.mean((0,1)) * hard.mean((0,1))).sum()
93 opt.zero_grad(); loss.backward(); opt.step()
94 net.eval()
95 with torch.no_grad(): pred, z, route = net(torch.as_tensor(ds['xte'], dtype=torch.float32, device=dev), balanced=balanced, shells=K)
96 mse = float(((pred - torch.as_tensor(ds['yte'], device=dev))**2).mean().cpu())
97 return {'metric': mse, 'routing': net.last_stats}
98
99def eval_cfg(cfg, seeds=SEEDS):
100 vals=[]; records=[]
101 for s in seeds:
102 ds=bench_get_dataset('correlated_token_moe_regression', s, 400, 200); r=train_one(ds,s,cfg['lr'],cfg.get('aux',0),cfg['balanced'],cfg.get('K',4))
103 vals.append(r['metric']); records.append({'seed':s,'metric':r['metric'],'routing':r['routing']})
104 return {'per_seed': vals, 'mean': float(np.mean(vals)), 'records': records, 'cfg': cfg}
105
106def main():
107 # Union of all learning rates is shared by baseline and idea; baseline's decisive aux knob is swept.
108 lrs=[0.0015,0.003,0.006]; auxs=[0.0,0.01,0.1]
109 base_grid=[{'lr':lr,'aux':a,'balanced':False,'K':4} for lr in lrs for a in auxs]
110 base_sweep=[]
111 for cfg in base_grid:
112 r=eval_cfg(cfg, SWEEP_SEEDS); base_sweep.append({'cfg':cfg,'mean':r['mean'],'per_seed':r['per_seed']})
113 best_cfg=min(base_sweep,key=lambda q:q['mean'])['cfg']
114 base_full=eval_cfg(best_cfg, SEEDS)
115 idea_grid=[{'lr':lr,'balanced':True,'K':k} for lr in lrs for k in [2,4,8]]
116 idea_trials=[]
117 for cfg in idea_grid:
118 r=eval_cfg(cfg, SEEDS); idea_trials.append(r)
119 idea=min(idea_trials,key=lambda q:q['mean'])
120 # Signature is measured on trained systems, not analytically assumed.
121 br=base_full['records']; ir=idea['records']
122 sig={'prediction':'per-shell expert count spread <= 1',
123 'baseline_mean_load_cv':float(np.mean([q['routing']['load_cv'] for q in br])),
124 'idea_mean_load_cv':float(np.mean([q['routing']['load_cv'] for q in ir])),
125 'baseline_mean_shell_spread':float(np.mean([q['routing']['max_shell_count_spread'] for q in br])),
126 'idea_mean_shell_spread':float(np.mean([q['routing']['max_shell_count_spread'] for q in ir])),
127 'idea_mean_reassigned_fraction':float(np.mean([q['routing']['reassigned_fraction'] for q in ir])),
128 'confirmed':bool(np.mean([q['routing']['max_shell_count_spread'] for q in ir]) <= 1.0)}
129 report=make_report('correlated_token_moe_regression','custom_moe',
130 {'best_cfg':best_cfg,'sweep':base_sweep,'full':base_full}, idea,
131 {'custom_track':{'name':META['name'],'file':'shell_moe_bench_track.py','domain':META['domain']}, **sig})
132 report['idea_trials']=[{'cfg':r['cfg'],'mean':r['mean'],'per_seed':r['per_seed']} for r in idea_trials]
133 (ROOT/'bench_report.json').write_text(json.dumps(report,indent=2))
134 print(json.dumps(report,indent=2))
135
136if __name__=='__main__': main()