import sys, json, math, time from pathlib import Path import numpy as np import torch from torch import nn sys.path.insert(0, '/home/maxwelhelp/all/math2nn') from bench import make_report, permutation_pvalue META = { 'name': 'heterogeneous_regime_moe', 'domain': 'moe-routing', 'description': 'Regression with heterogeneous input regimes requiring sparse expert routing.' } def get_dataset(seed, n_train=400, n_test=400): rng = np.random.default_rng(seed) def gen(n): x = rng.uniform(-1, 1, (n, 8)).astype('float32') z = np.argmax(x[:, :4], axis=1) y = np.zeros(n, dtype='float32') a = z == 0; y[a] = np.sin(3*x[a,4]) + .35*x[a,5]**2 a = z == 1; y[a] = x[a,4]*x[a,5] + .5*np.cos(2*x[a,6]) a = z == 2; y[a] = np.tanh(2*x[a,4]-x[a,6]) + .25*x[a,7] a = z == 3; y[a] = .6*x[a,4]**2 - .7*x[a,5] + np.sin(x[a,7]) y += rng.normal(0, .04, n).astype('float32') return x, y[:, None] xtr, ytr = gen(n_train); xte, yte = gen(n_test) return {'xtr': xtr, 'ytr': ytr, 'xte': xte, 'yte': yte, 'task': 'regression', 'metric': 'mse'} def seed_all(seed): np.random.seed(seed); torch.manual_seed(seed) if torch.cuda.is_available(): torch.cuda.manual_seed_all(seed) class SparseMoE(nn.Module): def __init__(self, alpha=0., beta=.9, aux=.05, topk=2, experts=8): super().__init__() self.E, self.topk = experts, topk self.alpha, self.beta, self.aux = alpha, beta, aux self.router = nn.Linear(8, experts) self.shared = nn.Sequential(nn.Linear(8, 32), nn.Tanh()) self.experts = nn.ModuleList([nn.Sequential(nn.Linear(32, 32), nn.Tanh(), nn.Linear(32, 1)) for _ in range(experts)]) self.register_buffer('pressure', torch.zeros(experts)) self.last_load = None self.last_raw_mass = None def forward(self, x, train_mode=True): h = self.shared(x) raw = self.router(x) soft = raw.softmax(-1) mass = soft.mean(0) if train_mode: self.pressure.mul_(self.beta).add_(mass.detach(), alpha=1-self.beta) if self.alpha != 0: lam = x.new_tensor([.25, 1., 4., 16.]) q = (torch.exp(-self.pressure[:, None] * lam[None, :]).mean(1)).clamp_min(1e-8) logits = raw + self.alpha * torch.log(q)[None, :] else: logits = raw vals, inds = logits.topk(self.topk, dim=-1) gates = vals.softmax(-1) out = torch.zeros(x.size(0), 1, device=x.device) load = torch.zeros(self.E, device=x.device) for k in range(self.topk): for e in range(self.E): mask = inds[:, k] == e if mask.any(): out[mask] += gates[mask, k:k+1] * self.experts[e](h[mask]) load[e] += mask.float().sum() self.last_load = (load / (x.size(0)*self.topk)).detach() self.last_raw_mass = mass.detach() # Standard Switch-style importance/load penalty for baseline only. aux_loss = self.E * (soft.mean(0) * (load / (x.size(0)*self.topk))).sum() return out, aux_loss def train_one(seed, cfg, idea): seed_all(seed) ds = get_dataset(seed, 200, 200) device = 'cuda' if torch.cuda.is_available() else 'cpu' try: model = SparseMoE(alpha=cfg['alpha'] if idea else 0., aux=cfg['aux'], beta=cfg['beta']).to(device) opt = torch.optim.Adam(model.parameters(), lr=cfg['lr'], weight_decay=cfg['wd']) x = torch.as_tensor(ds['xtr'], device=device); y = torch.as_tensor(ds['ytr'], device=device) xt = torch.as_tensor(ds['xte'], device=device); yt = torch.as_tensor(ds['yte'], device=device) rng = np.random.default_rng(seed+99); bs=64 for ep in range(cfg['epochs']): model.train(); order = rng.permutation(len(x)) for j in range(0, len(x), bs): ii = torch.as_tensor(order[j:j+bs], device=device) pred, aux = model(x[ii], True) loss = ((pred-y[ii])**2).mean() + (cfg['aux']*aux if not idea else 0.) opt.zero_grad(); loss.backward(); opt.step() model.eval() with torch.no_grad(): pred, _ = model(xt, False); mse=float(((pred-yt)**2).mean().cpu()) load=model.last_load.cpu().numpy(); raw=model.last_raw_mass.cpu().numpy() 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()} except Exception: # CPU fallback is explicit and deterministic. torch.cuda.empty_cache() if torch.cuda.is_available() else None if device == 'cuda': torch.cuda.is_available = lambda: False return train_one(seed, cfg, idea) raise def eval_cfg(cfg, idea, seeds=range(8)): vals=[]; sig=[] for s in seeds: v, z=train_one(int(s), cfg, idea); vals.append(v); sig.append(z) return {'mean':float(np.mean(vals)), 'std':float(np.std(vals)), 'per_seed':vals, 'n':len(vals), 'signatures':sig} def main(): # Union is shared: baseline is evaluated at every lr/central knob used by idea. 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),)] 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))] # Baseline method knob parity: auxiliary coefficient is swept, including zero. 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)] t=time.time(); baseline_sweep=[] for cfg in base_grid: r=eval_cfg(cfg,False,range(8)); baseline_sweep.append({'cfg':cfg,'mean':r['mean'],'per_seed':r['per_seed']}) best=min(baseline_sweep,key=lambda x:x['mean']); base_full=eval_cfg(best['cfg'],False,range(8)) idea_results=[] for cfg in idea_grid: r=eval_cfg(cfg,True,range(8)); idea_results.append((cfg,r)) best_i,best_r=min(idea_results,key=lambda cr:cr[1]['mean']) diffs=[i-b for i,b in zip(best_r['per_seed'],base_full['per_seed'])] 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])) 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)} report=make_report('heterogeneous_regime_moe','custom_sparse_moe',{'best_cfg':best['cfg'],'sweep':baseline_sweep,'full':base_full},best_r,extra) report['idea']['best_cfg']=best_i report['idea']['sweep']=[{'cfg':c,'mean':r['mean']} for c,r in idea_results] # Recompute comparison from raw evaluate() results; make_report receives raw per-seed results. report['comparison'] = __import__('bench', fromlist=['compare_results']).compare_results(base_full, best_r) report['custom_track']={'name':META['name'],'file':'stage2_bench.py','domain':META['domain']} report['comparison']['paired_delta_mean']=float(np.mean(diffs)); report['comparison']['permutation_pvalue']=permutation_pvalue(diffs) report['runtime_sec']=time.time()-t Path('bench_report.json').write_text(json.dumps(report,indent=2)) print(json.dumps(report,indent=2)) if __name__=='__main__': main()