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.