"""NumPy reference implementation of Laplace-heterogeneous MoE routing.""" import numpy as np class LaplaceHeterogeneousRouter: def __init__(self, n_experts, lambdas=(0.25, 1.0, 4.0, 16.0), weights=None, alpha=1.0, beta=0.9, epsilon=1e-12): self.n_experts = int(n_experts) self.lambdas = np.asarray(lambdas, dtype=float) if self.lambdas.ndim != 1 or len(self.lambdas) == 0 or np.any(self.lambdas < 0): raise ValueError("lambdas must be nonempty and nonnegative") if weights is None: weights = np.ones(len(self.lambdas)) / len(self.lambdas) self.weights = np.asarray(weights, dtype=float) if self.weights.shape != self.lambdas.shape or np.any(self.weights < 0) or self.weights.sum() <= 0: raise ValueError("weights must be nonnegative and match lambdas") self.weights /= self.weights.sum() self.alpha, self.beta, self.epsilon = float(alpha), float(beta), float(epsilon) self.pressure = np.zeros(self.n_experts) @staticmethod def softmax(logits): z = logits - logits.max(axis=1, keepdims=True) e = np.exp(z) return e / e.sum(axis=1, keepdims=True) def availability(self): return (self.weights[:, None] * np.exp(-self.lambdas[:, None] * self.pressure[None, :])).sum(axis=0) def effective_hazard(self): terms = self.weights[:, None] * np.exp(-self.lambdas[:, None] * self.pressure[None, :]) q = terms.sum(axis=0) return (terms * self.lambdas[:, None]).sum(axis=0) / np.maximum(q, self.epsilon) def route(self, logits, top_k=2): logits = np.asarray(logits, dtype=float) if logits.ndim != 2 or logits.shape[1] != self.n_experts: raise ValueError("logits must have shape [batch, n_experts]") if not 1 <= top_k <= self.n_experts: raise ValueError("invalid top_k") mass = self.softmax(logits).mean(axis=0) self.pressure = self.beta * self.pressure + (1.0 - self.beta) * mass q = self.availability() adjusted = logits + self.alpha * np.log(q + self.epsilon)[None, :] chosen = np.argpartition(-adjusted, top_k - 1, axis=1)[:, :top_k] return chosen, adjusted, q, self.pressure.copy()