🤖 AI Summary
This study addresses the gradient sparsity and suboptimal exploration-exploitation trade-off inherent in training Top-k selection modules for large language models. We propose a principled objective function that integrates supervised learning with policy gradients to generate hybrid optimization signals, establishing theoretical convergence bounds while balancing bias-variance trade-offs to overcome the limitations of heuristic selection. Within retrieval-augmented generation (RAG) and mixture-of-experts (MoE) architectures, the selection mechanism is jointly optimized via hybrid gradient descent and reinforcement learning strategies. Experimental results demonstrate that this approach significantly reduces perplexity and improves accuracy across synthetic regression, retrieval-augmented generation, and mixture-of-experts tasks.
📝 Abstract
Selecting the best $k$ items out of $m$ candidates is a critical component of modern large language model systems, such as document selection in Retrieval-Augmented Generation (RAG) and expert routing in Mixture-of-Experts (MoEs). However, training these selection modules remains challenging due to weak gradient signals and suboptimal exploration-exploitation tradeoffs. Furthermore, prior works often rely on heuristics, lacking principled objectives and approaches that explicitly model and solve the top-$k$ selection problem. In this work, we propose a principled objective for training selection modules, whose gradient naturally provides richer training signals in a hybrid form---containing a supervised-gradient component and a policy-gradient component. We show that the selection problem becomes harder as $m$ increases, and our algorithm converges at rate $O(1/\sqrt{T})$, with the optimal upper bound achieved by balancing between bias and variance. Practically, we apply our method to a set of tasks involving top-$k$ selection, including synthetic regression problems, RAG, and MoE systems, showing that our method outperforms the baselines in next-token prediction perplexity and QA accuracy.