import os, sys, json, itertools, copy import numpy as np import torch import torch.nn as nn ROOT = '/home/maxwelhelp/all/math2nn' sys.path.insert(0, ROOT) from bench import make_report, sweep_baseline, evaluate, get_dataset D, M = 6, 2 EPOCHS, BATCH = 22, 64 def stats(theta): subs = np.asarray(list(itertools.combinations(range(D), M)), dtype=int) a = theta[subs].sum(1); a -= a.max() p = np.exp(a); p /= p.sum() X = np.zeros((len(subs), D)); X[np.arange(len(subs))[:, None], subs] = 1 mu = p @ X cov = (X*p[:, None]).T @ X - np.outer(mu, mu) return mu, cov def math_check(): worst, bound_min, cap_err = 0., 1e9, 0. for a in np.linspace(0, 10, 9): th = np.zeros(D); th[0] = a; th[1] = -a/2 mu, s = stats(th); v = np.diag(s); V = v.sum() b = .5*(np.diag(v) - np.outer(v, v)/V) bound_min = min(bound_min, np.linalg.eigvalsh(s-b).min()) pinv = np.linalg.pinv(s, rcond=1e-11) for i in range(D): for j in range(i+1, D): q = np.zeros(D); q[i] = 1; q[j] = -1 worst = max(worst, (q@pinv@q)/(1/v[i]+1/v[j])) g = np.linspace(-1, 1, D); g -= g.mean() u = np.linalg.pinv(s, rcond=1e-11) @ g; u -= u.mean() raw = max(abs(u[i]-u[j])/np.sqrt(1/v[i]+1/v[j]) for i in range(D) for j in range(i+1,D)) rho=.23; scale=min(1.,rho/raw) if raw else 1. delta=u*scale obs=max(abs(delta[i]-delta[j])/np.sqrt(1/v[i]+1/v[j]) for i in range(D) for j in range(i+1,D)) cap_err=max(cap_err, abs(obs-min(rho,raw))) return {'max_resistance_ratio': float(worst), 'min_covariance_bound_eigenvalue': float(bound_min), 'trust_cap_max_abs_error': float(cap_err), 'passed': bool(worst <= 1+1e-8 and bound_min >= -1e-8 and cap_err < 1e-10)} class MoE(nn.Module): def __init__(self, temp=1.0): super().__init__(); self.temp=temp self.router=nn.Linear(4,D) self.experts=nn.ModuleList([nn.Sequential(nn.Linear(4,16),nn.Tanh(),nn.Linear(16,1)) for _ in range(D)]) def forward(self, x, idea=False): z=self.router(x)/self.temp outs=torch.cat([e(x) for e in self.experts], 1) if idea: # Exact fixed-m external-field inclusion means, computed by enumeration. combos=list(itertools.combinations(range(D),M)) scores=torch.stack([z[:,list(c)].sum(1) for c in combos],1) pp=torch.softmax(scores,1) gate=torch.zeros_like(z) for k,c in enumerate(combos): gate[:,list(c)] += pp[:,k:k+1] gate=gate/M else: gate=torch.softmax(z,1) return (gate*outs).sum(1,keepdim=True), z def run(seed, lr, idea, temp=1.0, rho=.3, return_sig=False): torch.manual_seed(seed); np.random.seed(seed) ds=get_dataset('router_regime_regression', seed, 400, 160) ds = {k: (torch.as_tensor(v, dtype=torch.float32) if k in ('xtr','ytr','xte','yte') else v) for k,v in ds.items()} model=MoE(temp=temp) device='cuda' if torch.cuda.is_available() else 'cpu' try: model=model.to(device); x=ds['xtr'].to(device); y=ds['ytr'].to(device) xt=ds['xte'].to(device); yt=ds['yte'].to(device) opt=torch.optim.Adam(model.parameters(),lr=lr) lossf=nn.MSELoss(); cap_obs=[]; cap_pred=[]; raw_vals=[] for ep in range(EPOCHS): model.train(); perm=torch.randperm(len(x),device=device) for ii in range(0,len(x),BATCH): q=perm[ii:ii+BATCH]; pred,z=model(x[q],idea); loss=lossf(pred,y[q]) opt.zero_grad(); loss.backward() if idea: # Router gradients are preconditioned in logit coordinates by Sigma^dagger. with torch.no_grad(): zz=z.detach().mean(0).cpu().numpy(); mu,s=stats(zz) v=np.clip(np.diag(s),1e-7,None); pinv=np.linalg.pinv(s,rcond=1e-10) P=np.eye(D)-np.ones((D,D))/D for p in [model.router.weight, model.router.bias]: if p.grad is None: continue gg=p.grad.detach().cpu().numpy(); gg=P@gg if gg.ndim==1 else (P@gg) uu=pinv@gg; uu=P@uu raw=float(np.max([np.max(np.abs(uu[i]-uu[j])/np.sqrt(1/v[i]+1/v[j])) for i in range(D) for j in range(i+1,D)])) scale=min(1.,rho/raw) if raw>0 else 1. obs=raw*scale; raw_vals.append(raw); cap_obs.append(obs); cap_pred.append(min(rho,raw)) p.grad.copy_(torch.as_tensor(uu*scale,dtype=p.grad.dtype,device=device)) opt.step() model.eval() with torch.no_grad(): metric=float(lossf(model(xt,idea)[0],yt).cpu()) sig={'mean_observed_cap':float(np.mean(cap_obs)) if cap_obs else None, 'mean_predicted_cap':float(np.mean(cap_pred)) if cap_pred else None, 'max_observed_cap':float(np.max(cap_obs)) if cap_obs else None, 'rho':rho, 'n_trained_updates':len(cap_obs)} if return_sig: return metric,sig return metric except RuntimeError: if device=='cuda': torch.cuda.empty_cache(); torch.set_default_device('cpu') return run(seed,lr,idea,temp,rho,return_sig) raise def baseline_factory(cfg): return lambda seed: run(seed, cfg['lr'], False, cfg['temp']) def main(): check=math_check(); assert check['passed'], check # Baseline sweep includes every idea learning rate (search-space parity) and its temperature knob. grid=[{'lr':lr,'temp':temp} for lr in (.003,.006,.012) for temp in (.7,1.0)] base=sweep_baseline(baseline_factory, grid, seeds=(0,1,2,3)) best=base['best_cfg']; lrs=[.003,.006,.012] near=[x for x in lrs if x != best['lr']][:2] idea_cfgs=[{'lr':best['lr'],'temp':best['temp'],'rho':r} for r in (.15,.30,.60)] # nearby learning rates are also evaluated, and all occur in the baseline grid. idea_cfgs += [{'lr':lr,'temp':best['temp'],'rho':.30} for lr in near] idea_runs=[]; best_run=None for cfg in idea_cfgs: vals=[]; sigs=[] for seed in range(8): v,s=run(seed,cfg['lr'],True,cfg['temp'],cfg['rho'],True); vals.append(v); sigs.append(s) r={'cfg':cfg,'per_seed':vals,'mean':float(np.mean(vals)),'std':float(np.std(vals,ddof=1)), 'signature_summary':{'mean_observed_cap':float(np.mean([s['mean_observed_cap'] for s in sigs])), 'mean_predicted_cap':float(np.mean([s['mean_predicted_cap'] for s in sigs]))}} idea_runs.append(r) if best_run is None or r['mean']