Laplace-Heterogeneous MoE Routing / stage2_bench.py

Failed on benchmark

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