Demonstration selection made efficient using state space models
Long-Context Demonstration Selection Using State Space Models
Machine LearningComputation and Language
Summary
Choosing examples to show a language model before asking it a question can be slow and costly when the examples are long. The authors developed a method that uses simpler mathematical models called state space models to imitate parts of the original complex model more quickly. This lets them pick helpful examples faster and use less computing power without losing accuracy. Their tests showed this method works well on text classification and reasoning tasks, making the process both faster and more accurate.
What this means in practice
- •For natural language processing teams: Reduce computational cost when selecting example prompts for language models in applications requiring long inputs.
- •For software engineers building ai assistants: Speed up and improve accuracy of demonstration selection during model inference for text comprehension and reasoning tasks.
Authors
Ziniu Zhang, Zhenshuo Zhang, Ruoxuan Xiong, Gene Cooperman, Hongyang R. Zhang
Abstract
We study the problem of demonstration selection, which involves selecting a subset of examples for prepending to a query to a language model. This problem is closely related to in-context learning and language model inference. Since the inference cost of a transformer model scales quadratically with sequence length, the selection problem becomes especially challenging in a long-context scenario. In this paper, we tackle this problem by building on state space models (SSMs), which require only linear inference time given the input. Our approach involves two algorithms. The first learns a small set of SSMs through distillation of a (trained) transformer model. We partition all the layers into consecutive groups. Then for each group, we estimate a separate state space model to replicate the input-output behavior within the adjacent layers. Second, we map the distilled model outputs to a small set of tokens, and apply these embeddings for demonstration selection in downstream applications. We perform extensive experiments in both synthetic and real-world datasets to validate our approach. We demonstrate that the distilled SSMs only incur an approximation error of less than $0.7\%$ relative to the true output. In downstream evaluation, we show that on several text classification and reasoning tasks, our approach reduces FLOPs by $14.2\times$ and improves accuracy by $6.48\%$ relative to baseline demonstration selection methods.