Effective-resistance natural-gradient routing / router.py
Mechanism confirmed, baseline not beaten
1import itertools
2import numpy as np
3
4
5def exact_stats(theta, m):
6 """Exact external-field fixed-m marginals and covariance for small d."""
7 theta = np.asarray(theta, dtype=float)
8 d = theta.size
9 subsets = np.array(list(itertools.combinations(range(d), m)), dtype=int)
10 scores = theta[subsets].sum(axis=1)
11 scores -= scores.max()
12 p = np.exp(scores)
13 p /= p.sum()
14 X = np.zeros((len(subsets), d))
15 X[np.arange(len(subsets))[:, None], subsets] = 1.0
16 mu = p @ X
17 sigma = (X * p[:, None]).T @ X - np.outer(mu, mu)
18 return mu, sigma
19
20
21def sample_subset_dp(theta, m, rng=None):
22 """Exact O(d*m) backward sampler from exp(sum(theta_i)) over |S|=m."""
23 theta = np.asarray(theta, dtype=float)
24 d = theta.size
25 if not 0 <= m <= d:
26 raise ValueError("m must be between 0 and d")
27 rng = np.random.default_rng() if rng is None else rng
28 w = np.exp(theta - np.max(theta))
29 E = np.zeros((d + 1, m + 1), dtype=float)
30 E[0, 0] = 1.0
31 for k in range(1, d + 1):
32 E[k] = E[k - 1]
33 for r in range(1, min(k, m) + 1):
34 E[k, r] += w[k - 1] * E[k - 1, r - 1]
35 chosen = []
36 remaining = m
37 for k in range(d, 0, -1):
38 if remaining == 0:
39 break
40 total = E[k, remaining]
41 include = w[k - 1] * E[k - 1, remaining - 1]
42 if rng.random() < include / total:
43 chosen.append(k - 1)
44 remaining -= 1
45 if remaining != 0:
46 raise RuntimeError("DP sampler failed to select m experts")
47 return np.array(sorted(chosen), dtype=int)
48
49
50def natural_direction(theta, m, grad, rho=0.3, damping=0.0):
51 """Sigma-dagger projected gradient with pairwise resistance trust cap."""
52 theta = np.asarray(theta, dtype=float)
53 grad = np.asarray(grad, dtype=float)
54 d = theta.size
55 mu, sigma = exact_stats(theta, m)
56 P = np.eye(d) - np.ones((d, d)) / d
57 projected = P @ grad
58 if damping:
59 direction = np.linalg.solve(sigma + damping * P + np.ones((d, d)) / d,
60 projected)
61 else:
62 direction = np.linalg.pinv(sigma, rcond=1e-11) @ projected
63 direction = P @ direction
64 v = np.diag(sigma)
65 normalized = [abs(direction[i] - direction[j]) /
66 np.sqrt(1.0 / v[i] + 1.0 / v[j])
67 for i in range(d) for j in range(i + 1, d)]
68 raw = max(normalized, default=0.0)
69 scale = min(1.0, rho / raw) if raw > 0 else 1.0
70 return direction * scale, {"mu": mu, "sigma": sigma, "raw_cap": raw,
71 "scale": scale}