Effective-resistance natural-gradient routing / router.py

Mechanism confirmed, baseline not beaten

Raw ⬇ ZIP
 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}