import json from pathlib import Path import numpy as np SEED = 1234 def shell_ids(conf, K): order = np.argsort(conf, kind='stable') out = np.empty(len(conf), dtype=np.int64) for j, i in enumerate(order): out[i] = min(K - 1, (j * K) // len(conf)) return out def shell_balance_route(logits, shells, K): """Top-1 exact per-shell balanced routing via greedy maximum-weight assignment.""" n, e = logits.shape route = np.empty(n, dtype=np.int64) reassigned = 0 score_loss = 0.0 for k in range(K): ids = np.flatnonzero(shells == k) b = len(ids) if b == 0: continue target = np.full(e, b // e, dtype=int) rem = b - target.sum() demand = np.bincount(np.argmax(logits[ids], axis=1), minlength=e) for x in np.argsort(-demand, kind='stable')[:rem]: target[x] += 1 pairs = [(float(logits[i, x]), int(i), x) for i in ids for x in range(e)] pairs.sort(key=lambda t: (-t[0], t[1], t[2])) used = np.zeros(e, dtype=int) assigned = set() for val, i, x in pairs: if i not in assigned and used[x] < target[x]: route[i] = x used[x] += 1 assigned.add(i) assert len(assigned) == b and np.array_equal(used, target) pref = np.argmax(logits[ids], axis=1) reassigned += int(np.sum(route[ids] != pref)) score_loss += float(np.sum(logits[ids, pref] - logits[ids, route[ids]])) return route, reassigned / n, score_loss / n def cv(route, E): loads = np.bincount(route, minlength=E).astype(float) return float(loads.std() / loads.mean()), loads.tolist() def toy_predictions(): rng = np.random.default_rng(SEED) rows = [] # Eq. (66) prediction: quota counts differ by at most one in every shell, # for all tested expert counts, shell counts, and batch sizes. for E in [3, 5, 8]: for K in [2, 4, 7]: n = 997 logits = rng.normal(size=(n, E)) sh = shell_ids(np.max(logits, axis=1), K) route, _, _ = shell_balance_route(logits, sh, K) worst = 0 for k in range(K): counts = np.bincount(route[sh == k], minlength=E) worst = max(worst, int(counts.max() - counts.min())) rows.append({'prediction':'per-shell count spread <= 1', 'E':E, 'K':K, 'observed':worst, 'predicted_max':1}) # Eq. (66) also predicts identical polynomial marginals. For each shell, # evaluate p_e(P)=sum_k count[e,k] P^k(1-P)^(K-k) on a probability grid. for E, K in [(4, 3), (8, 5), (8, 8)]: n = E * K * 32 logits = rng.normal(size=(n, E)) sh = shell_ids(np.max(logits, axis=1), K) route, _, _ = shell_balance_route(logits, sh, K) coeff = np.zeros((E, K)) base_coeff = np.zeros((E, K)) base = np.argmax(logits, axis=1) for k in range(K): for e in range(E): coeff[e, k] = np.sum((route == e) & (sh == k)) base_coeff[e, k] = np.sum((base == e) & (sh == k)) ps = np.linspace(0.01, 0.99, 99) idea_spread = [] base_spread = [] for P in ps: q = np.array([P**k * (1-P)**(K-k) for k in range(K)]) idea_spread.append(float(np.ptp(coeff @ q))) base_spread.append(float(np.ptp(base_coeff @ q))) rows.append({'prediction':'Eq.66 marginal spread is zero for all P', 'E':E, 'K':K, 'observed_max_spread':max(idea_spread), 'predicted':0.0, 'baseline_max_spread':max(base_spread)}) return rows def comparison(): rng = np.random.default_rng(SEED + 1) E, K, n = 8, 8, 4096 # Correlated logits intentionally create expert monopoly in confidence shells. common = rng.normal(size=(n, 1)) logits = rng.normal(scale=.45, size=(n, E)) logits[:, 0] += 1.8 * np.maximum(common[:, 0], 0) logits[:, 1] += 1.1 * np.maximum(-common[:, 0], 0) conf = np.max(logits, axis=1) shells = shell_ids(conf, K) base = np.argmax(logits, axis=1) idea, frac, loss = shell_balance_route(logits, shells, K) bcv, bloads = cv(base, E) icv, iloads = cv(idea, E) shell_cvs = [] for k in range(K): ids = np.flatnonzero(shells == k) shell_cvs.append({'shell':k, 'baseline_cv':cv(base[ids], E)[0], 'idea_cv':cv(idea[ids], E)[0], 'size':len(ids)}) return {'baseline_cv':bcv, 'idea_cv':icv, 'baseline_loads':bloads, 'idea_loads':iloads, 'reassigned_fraction':frac, 'router_score_loss_per_token':loss, 'shells':shell_cvs} def main(): out = {'seed':SEED, 'toy_predictions':toy_predictions(), 'comparison':comparison()} Path('results.json').write_text(json.dumps(out, indent=2)) print(json.dumps(out, indent=2)) if __name__ == '__main__': main()