Shell-Wise Balanced MoE Routing / stage2_shell_moe.py

Mechanism confirmed, baseline not beaten

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