import itertools import numpy as np def exact_stats(theta, m): """Exact external-field fixed-m marginals and covariance for small d.""" theta = np.asarray(theta, dtype=float) d = theta.size subsets = np.array(list(itertools.combinations(range(d), m)), dtype=int) scores = theta[subsets].sum(axis=1) scores -= scores.max() p = np.exp(scores) p /= p.sum() X = np.zeros((len(subsets), d)) X[np.arange(len(subsets))[:, None], subsets] = 1.0 mu = p @ X sigma = (X * p[:, None]).T @ X - np.outer(mu, mu) return mu, sigma def sample_subset_dp(theta, m, rng=None): """Exact O(d*m) backward sampler from exp(sum(theta_i)) over |S|=m.""" theta = np.asarray(theta, dtype=float) d = theta.size if not 0 <= m <= d: raise ValueError("m must be between 0 and d") rng = np.random.default_rng() if rng is None else rng w = np.exp(theta - np.max(theta)) E = np.zeros((d + 1, m + 1), dtype=float) E[0, 0] = 1.0 for k in range(1, d + 1): E[k] = E[k - 1] for r in range(1, min(k, m) + 1): E[k, r] += w[k - 1] * E[k - 1, r - 1] chosen = [] remaining = m for k in range(d, 0, -1): if remaining == 0: break total = E[k, remaining] include = w[k - 1] * E[k - 1, remaining - 1] if rng.random() < include / total: chosen.append(k - 1) remaining -= 1 if remaining != 0: raise RuntimeError("DP sampler failed to select m experts") return np.array(sorted(chosen), dtype=int) def natural_direction(theta, m, grad, rho=0.3, damping=0.0): """Sigma-dagger projected gradient with pairwise resistance trust cap.""" theta = np.asarray(theta, dtype=float) grad = np.asarray(grad, dtype=float) d = theta.size mu, sigma = exact_stats(theta, m) P = np.eye(d) - np.ones((d, d)) / d projected = P @ grad if damping: direction = np.linalg.solve(sigma + damping * P + np.ones((d, d)) / d, projected) else: direction = np.linalg.pinv(sigma, rcond=1e-11) @ projected direction = P @ direction v = np.diag(sigma) normalized = [abs(direction[i] - direction[j]) / np.sqrt(1.0 / v[i] + 1.0 / v[j]) for i in range(d) for j in range(i + 1, d)] raw = max(normalized, default=0.0) scale = min(1.0, rho / raw) if raw > 0 else 1.0 return direction * scale, {"mu": mu, "sigma": sigma, "raw_cap": raw, "scale": scale}