import sys, json, random from pathlib import Path import numpy as np import torch import torch.nn as nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import make_report, get_dataset as bench_get_dataset from shell_moe_bench_track import META ROOT = Path(__file__).parent SEEDS = tuple(range(8)) SWEEP_SEEDS = tuple(range(4)) E, TOK, D = 8, 16, 4 def seed_all(seed): random.seed(seed); np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) def device(): return torch.device('cuda' if torch.cuda.is_available() else 'cpu') class MoE(nn.Module): def __init__(self, hidden=32): super().__init__() self.router = nn.Sequential(nn.Linear(D, hidden), nn.Tanh(), nn.Linear(hidden, E)) self.experts = nn.ModuleList([nn.Sequential(nn.Linear(D, hidden), nn.ReLU(), nn.Linear(hidden, 1)) for _ in range(E)]) self.last_stats = {} def forward(self, x, balanced=False, shells=4): # x: batch, tokens, features. Return mean token output and routing logits/mask. z = self.router(x) pref = z.argmax(-1) if balanced: route = shell_route(z.detach().cpu().numpy(), shells) route = torch.as_tensor(route, device=x.device, dtype=torch.long) else: route = pref outs = torch.stack([m(x)[:, :, 0] for m in self.experts], dim=-1) chosen = outs.gather(-1, route.unsqueeze(-1)).squeeze(-1) self.last_stats = routing_stats(z.detach(), route.detach(), pref.detach(), shells) return chosen.mean(1, keepdim=True), z, route def shell_route(logits, K): b, t, e = logits.shape flat = logits.reshape(-1, e) conf = flat.max(1) order = np.argsort(conf, kind='stable') sh = np.empty(len(conf), dtype=np.int64) for j, i in enumerate(order): sh[i] = min(K-1, j*K//len(conf)) route = np.empty(len(conf), dtype=np.int64) for k in range(K): ids = np.flatnonzero(sh == k); n = len(ids) if not n: continue target = np.full(e, n//e, dtype=int) # distribute remainders to experts with greatest preferred demand pref = flat[ids].argmax(1) demand = np.bincount(pref, minlength=e) for q in np.argsort(-demand, kind='stable')[:n % e]: target[q] += 1 # greedy maximum-weight quota assignment; each token gets one expert pairs = sorted(((float(flat[i, q]), i, q) for i in ids for q in range(e)), reverse=True) used = np.zeros(e, dtype=int); assigned = set() for val, i, q in pairs: if i not in assigned and used[q] < target[q]: route[i] = q; used[q] += 1; assigned.add(i) assert len(assigned) == n return route.reshape(b, t) def routing_stats(z, route, pref, K): a = z.reshape(-1, E).cpu().numpy(); r = route.reshape(-1).cpu().numpy(); p = pref.reshape(-1).cpu().numpy() conf = a.max(1); order = np.argsort(conf, kind='stable'); sh = np.empty(len(conf), int) for j, i in enumerate(order): sh[i] = min(K-1, j*K//len(conf)) loads = np.bincount(r, minlength=E).astype(float); cv = float(loads.std()/loads.mean()) spreads = [] for k in range(K): q = np.bincount(r[sh == k], minlength=E) if len(q): spreads.append(int(q.max()-q.min())) return {'load_cv': cv, 'loads': loads.tolist(), 'max_shell_count_spread': max(spreads or [0]), 'reassigned_fraction': float(np.mean(r != p)), 'shells': K} def train_one(ds, seed, lr, aux, balanced, K, epochs=12): seed_all(seed); dev = device() net = MoE().to(dev); opt = torch.optim.Adam(net.parameters(), lr=lr) x = torch.as_tensor(ds['xtr'], dtype=torch.float32, device=dev); y = torch.as_tensor(ds['ytr'], dtype=torch.float32, device=dev) for ep in range(epochs): net.train(); perm = torch.randperm(len(x), device=dev) for ii in range(0, len(x), 128): ix = perm[ii:ii+128]; pred, z, route = net(x[ix], balanced=balanced, shells=K) loss = ((pred-y[ix])**2).mean() if not balanced and aux > 0: prob = z.softmax(-1); hard = torch.nn.functional.one_hot(route, E).float() # standard Switch-style importance/load auxiliary loss loss = loss + aux * E * (prob.mean((0,1)) * hard.mean((0,1))).sum() opt.zero_grad(); loss.backward(); opt.step() net.eval() with torch.no_grad(): pred, z, route = net(torch.as_tensor(ds['xte'], dtype=torch.float32, device=dev), balanced=balanced, shells=K) mse = float(((pred - torch.as_tensor(ds['yte'], device=dev))**2).mean().cpu()) return {'metric': mse, 'routing': net.last_stats} def eval_cfg(cfg, seeds=SEEDS): vals=[]; records=[] for s in seeds: 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)) vals.append(r['metric']); records.append({'seed':s,'metric':r['metric'],'routing':r['routing']}) return {'per_seed': vals, 'mean': float(np.mean(vals)), 'records': records, 'cfg': cfg} def main(): # Union of all learning rates is shared by baseline and idea; baseline's decisive aux knob is swept. lrs=[0.0015,0.003,0.006]; auxs=[0.0,0.01,0.1] base_grid=[{'lr':lr,'aux':a,'balanced':False,'K':4} for lr in lrs for a in auxs] base_sweep=[] for cfg in base_grid: r=eval_cfg(cfg, SWEEP_SEEDS); base_sweep.append({'cfg':cfg,'mean':r['mean'],'per_seed':r['per_seed']}) best_cfg=min(base_sweep,key=lambda q:q['mean'])['cfg'] base_full=eval_cfg(best_cfg, SEEDS) idea_grid=[{'lr':lr,'balanced':True,'K':k} for lr in lrs for k in [2,4,8]] idea_trials=[] for cfg in idea_grid: r=eval_cfg(cfg, SEEDS); idea_trials.append(r) idea=min(idea_trials,key=lambda q:q['mean']) # Signature is measured on trained systems, not analytically assumed. br=base_full['records']; ir=idea['records'] sig={'prediction':'per-shell expert count spread <= 1', 'baseline_mean_load_cv':float(np.mean([q['routing']['load_cv'] for q in br])), 'idea_mean_load_cv':float(np.mean([q['routing']['load_cv'] for q in ir])), 'baseline_mean_shell_spread':float(np.mean([q['routing']['max_shell_count_spread'] for q in br])), 'idea_mean_shell_spread':float(np.mean([q['routing']['max_shell_count_spread'] for q in ir])), 'idea_mean_reassigned_fraction':float(np.mean([q['routing']['reassigned_fraction'] for q in ir])), 'confirmed':bool(np.mean([q['routing']['max_shell_count_spread'] for q in ir]) <= 1.0)} report=make_report('correlated_token_moe_regression','custom_moe', {'best_cfg':best_cfg,'sweep':base_sweep,'full':base_full}, idea, {'custom_track':{'name':META['name'],'file':'shell_moe_bench_track.py','domain':META['domain']}, **sig}) report['idea_trials']=[{'cfg':r['cfg'],'mean':r['mean'],'per_seed':r['per_seed']} for r in idea_trials] (ROOT/'bench_report.json').write_text(json.dumps(report,indent=2)) print(json.dumps(report,indent=2)) if __name__=='__main__': main()