基于状态空间模型的长上下文示范选择

arXiv cs.LG 论文

摘要

本文提出使用状态空间模型来高效选择用于长上下文语言模型提示的示范,从而降低计算成本并提升性能。

arXiv:2609.17888v1 公告类型:新 摘要:我们研究示范选择问题,即选择一个示例子集,将其前置到语言模型的查询中。该问题与上下文学习和语言模型推理密切相关。由于Transformer模型的推理成本随序列长度二次方增长,在长上下文场景下选择问题变得尤其具有挑战性。在本文中,我们通过状态空间模型(SSMs)来解决这一问题,SSMs在给定输入时仅需要线性推理时间。我们的方法涉及两个算法。第一个算法通过蒸馏(训练后的)Transformer模型来学习一组小的SSMs。我们将所有层划分为连续的组,然后为每个组估计一个独立的状态空间模型,以复制相邻层内的输入-输出行为。其次,我们将蒸馏模型的输出映射到一组小的令牌,并将这些嵌入应用于下游应用中的示范选择。我们在合成数据和真实世界数据集上进行了广泛实验以验证我们的方法。我们证明蒸馏SSMs相对于真实输出仅产生低于$0.7\%$的近似误差。在下游评估中,我们展示在多个文本分类和推理任务上,我们的方法相对于基线示范选择方法减少了$14.2\times$ FLOPs并提升了$6.48\%$的准确率。
查看原文
查看缓存全文

缓存时间: 2026/09/17 09:00

# 基于状态空间模型的长上下文示例选择
来源:https://arxiv.org/html/2609.17888

###### 摘要
我们研究示例选择问题,即为语言模型查询预置一组示例子集。该问题与上下文学习及语言模型推理密切相关。由于Transformer模型的推理成本随序列长度呈二次方增长,在长上下文场景下该选择问题尤为具有挑战性。本文基于状态空间模型(SSM)解决此问题,SSM在给定输入下仅需线性推理时间。我们的方法包含两个算法:首先通过**蒸馏**(训练后的)Transformer模型来学习一组小规模SSM。我们将所有层划分为连续组,针对每组估计独立的状态空间模型以模拟相邻层的输入输出行为。其次,将蒸馏模型的输出映射至一组小规模词元,并利用这些嵌入向量在下游应用中进行示例选择。我们在合成数据与真实数据集上进行了广泛实验验证。结果表明蒸馏后的SSM相对真实输出仅产生低于0.7%的近似误差。在下游评估中,针对多项文本分类与推理任务,我们的方法相比基线示例选择方法实现了14.2倍的FLOPs缩减与6.48%的准确率提升。

## 1 引言
语言模型在推理时日益依赖于处理长上下文提示(Oncescu et al., 2025 (https://arxiv.org/html/2609.17888#bib.bib7))。其中一种方法是基于历史轨迹中的示例对模型进行条件化(Xiong, 2025 (https://arxiv.org/html/2609.17888#bib.bib22))。具体而言,给定一个查询,目标是从规模为 n 的大型候选池中选取包含 k 个示例的子集。该问题被称为**示例选择**,与上下文学习密切相关(Garg et al., 2022 (https://arxiv.org/html/2609.17888#bib.bib1); Zhang et al., 2025b (https://arxiv.org/html/2609.17888#bib.bib2))。

图1:我们的方法概述:给定来自未知分布的候选示例集及查询-目标对,我们的方法通过从预训练模型蒸馏来学习一组状态空间模型。随后,利用从状态空间模型中提取的嵌入向量进行示例选择。我们采样多个示例子集并评估拼接后的子集嵌入,将每个示例的子集损失聚合成亲和力得分,最终选择得分最高的前 k 个示例。

示例选择的挑战在于单个示例的效果往往取决于与其同时出现在提示中的其他示例。评估候选子集通常需要针对多种不同提示反复运行语言模型(Zhang et al., 2025b (https://arxiv.org/html/2609.17888#bib.bib2))。对于Transformer模型,处理长提示的成本随总上下文长度 T 呈二次方增长,而 T 同时取决于示例的数量与长度。这导致在 k 增大时穷举评估众多候选子集变得不可行,从而形成计算瓶颈。

该挑战将示例选择与更广泛的长上下文高效推理问题联系起来。现有方法主要分为三类:第一,状态空间模型(SSM)(Gu et al., 2020 (https://arxiv.org/html/2609.17888#bib.bib5); Gu et al., 2022 (https://arxiv.org/html/2609.17888#bib.bib4); Gu and Dao, 2024 (https://arxiv.org/html/2609.17888#bib.bib9))及相关混合架构(Ren et al., 2025 (https://arxiv.org/html/2609.17888#bib.bib8); Oncescu et al., 2025 (https://arxiv.org/html/2609.17888#bib.bib7))实现了高效序列建模,但通常需要对底层语言模型进行重训。第二,免训练方法通过将注意力限制在局部窗口或选定词元上来加速推理(Xiao et al., 2024 (https://arxiv.org/html/2609.17888#bib.bib6); Xiao et al., 2025a (https://arxiv.org/html/2609.17888#bib.bib12))。第三,其他相关方法(Mu et al., 2023 (https://arxiv.org/html/2609.17888#bib.bib15); Chevalier et al., 2023 (https://arxiv.org/html/2609.17888#bib.bib16); Ge et al., 2024 (https://arxiv.org/html/2609.17888#bib.bib14))仍需处理完整提示。

本文基于状态空间模型扩展推理时的长上下文示例选择。核心思想是训练一组SSM,将每个候选子集的拼接示例映射为一组小规模词元。这些词元仅需对每个子集计算一次,即可在评估该子集的所有查询中重复使用,从而大幅降低评估多个候选示例子集的成本。

我们的算法包含三个主要组件:首先通过蒸馏学习小规模状态空间模型。将Transformer层划分为若干连续层的互斥组,为每组分配一个SSM。每个SSM将长示例提示的嵌入向量扫描为紧凑的隐藏表示,再映射至一组键值状态。输出的键值状态组合后形成Transformer可直接处理的词元,这些词元在推理时替代原始长提示。其次,我们设计了基于子集采样的示例选择算法(Li et al., 2023b (https://arxiv.org/html/2609.17888#bib.bib21); Li et al., 2023a (https://arxiv.org/html/2609.17888#bib.bib28); Li et al., 2024a (https://arxiv.org/html/2609.17888#bib.bib18); Li et al., 2024b (https://arxiv.org/html/2609.17888#bib.bib19); Zhang et al., 2025b (https://arxiv.org/html/2609.17888#bib.bib2))。关键在于该算法现基于(蒸馏后的)状态空间模型运行。我们反复采样候选子集,利用所提方法生成的表示评估其预测性能,并基于包含该示例的子集表现来估计每个示例的贡献。

最后,我们对所提框架进行实证分析。实验表明,该方法以低于0.6%的相对误差保持了原始模型的预测分布。在分类与推理基准测试中,我们的方法最高可实现14.2倍的FLOPs缩减,并较基线提升6.48%的下游准确率。

总之,我们设计了可扩展至长上下文推理的示例选择算法:首先构建将长示例提示映射为少量可复用词元的SSM;其次开发基于亲和力的子集选择算法,利用该表示评估众多候选示例子集;最后在合成数据、分类与推理基准测试中验证了算法有效性。实验代码已开源:https://github.com/VirtuosoResearch/Long-context-demonstration-selection

## 2 预备知识
我们研究构建提示过程中的示例选择。令 D 表示规模为 n 的候选示例集。令 (q,y) ∼ P 表示来自未知分布的查询-目标对。针对每个查询,我们选择包含 k 个示例且具有固定顺序的子集 D′⊆D。给定模型 f_W 及在输入域 X 和输出域 Y 上定义的损失函数 ℓ:X×Y→ℝ,**示例(子集)选择问题**定义为以下最小化问题:
$$\min_{\begin{subarray}{c}D^{\prime}\subseteq D:\\,\left\lvert D^{\prime}\right\rvert=k\end{subarray}}\mathbb{E}_{(q,y)\sim\mathcal{P}}\left[\ell\!\left(f_{W}(D^{\prime},q),y\right)\right].$$

对于选定的示例子集 D′,我们将它们拼接作为输入示例提示。令 X=[x₁,…,x_T]ᵀ∈ℝ^{T×H} 表示提示的词元嵌入向量,其中 T 为词元总数,H 为嵌入维度。由于 T 随示例数量和长度增长,长上下文提示可能产生非常长的输入序列。在Transformer模型中,处理提示的计算成本随 T 呈二次方增长,因为自注意力机制需计算所有词元间的成对交互。因此,评估大 T 下的推理结果计算代价高昂。

处理长提示的自然方式是采用状态空间模型(SSM),其通过循环更新总结输入序列。给定提示嵌入 X,离散SSM在每个词元位置 t 更新其状态:
$$h_{t}=\bar{A}h_{t-1}+\bar{B}x_{t}, \tag{1}$$
其中 h_t∈ℝ^N 为 N 维状态,Ā∈ℝ^{N×N} 为状态转移矩阵,B̄∈ℝ^{N×H} 将词元嵌入 x_t 映射至状态空间。SSM 仅处理每个词元一次,具有关于序列长度 T 的线性复杂度(Gu et al., 2022 (https://arxiv.org/html/2609.17888#bib.bib4))。先前工作表明,HiPPO-LegS矩阵通过维护输入历史的在线多项式投影来构建结构化的状态转移矩阵(Gu et al., 2020 (https://arxiv.org/html/2609.17888#bib.bib5); Gu and Dao, 2024 (https://arxiv.org/html/2609.17888#bib.bib9))。该状态近似于缩放勒让德基下的历史输入系数。低阶系数捕获序列历史的主要结构,高阶系数表示更细微的变化。因此,有限维HiPPO状态提供了长输入序列的紧凑摘要。

在上下文学习设定中,我们假设大型示例集中与任务相关的信息可通过HiPPO矩阵的低阶勒让德基分量得以保留。SSM 对完整输入仅需扫描一次,因此可替代Transformer模型对示例的推理,将推理复杂度从 O(T²) 降至 O(T)。

基于上述成果,自然引出一个问题:能否调整SSM来处理示例子集选择问题?
- 第一,如何为上下文学习设定设计SSM?
- 第二,如何利用SSM进行示例选择?

下节将分别设计算法回答上述问题。

## 3 我们的方法
我们构建SSM以实现高效的长上下文推理。首先,我们介绍SSM架构,并提供支持假设的理论与实证依据,同时引入训练SSM的蒸馏流程。随后,我们提出基于嵌入的示例选择方法,并通过受控合成实验评估其亲和力估计效果。

### 3.1 学习状态空间模型
我们设计采用HiPPO-LegS矩阵的SSM,以映射长示例提示并保留序列的主导低阶结构。在Transformer中,提示首先被处理为层级键值(KV)缓存,其中每层Transformer存储所有提示词元的键和值表示。这些缓存表示在解码阶段使用。我们并非计算完整提示的确切KV缓存,而是使用SSM处理它。SSM 顺序扫描提示,将长序列总结为紧凑的隐藏表示。

我们将Transformer模型 f 的 m 层划分为 g 个互斥的连续层组,为每组生成低维目标。每组分配独立的SSM,仅预测其对应层组的KV状态。具体而言,对于每个层组 G_i(i=1,…,g),SSM处理提示嵌入 X 并生成最终隐藏状态 h_T^{(i)}∈ℝ^N。随后多层感知机 f_W^{(i)} 将 h_T^{(i)} 投影为组内每个Transformer层的独立键值状态对:
$$\left\{\bigl(\tilde{K}^{(l)},\tilde{V}^{(l)}\bigr)\right\}_{l\in G_{i}}=f_{W}^{(i)}(h_{T}^{(i)}(X)).$$
对于每个Transformer层 l,原始KV状态维度为 T×n_kv×d,而投影后状态维度为 n_v×n_kv×d,其中 n_kv 为键值头数量,d 为每头维度。所有 g 个SSM输出生成的KV状态对齐至相同的 n_v 词元位置,并跨层组组装形成这 n_v 个词元的层级KV缓存。因此,投影在保持键值头和头维度轴的同时,将 T 个词元映射为 n_v 个词元供Transformer直接处理。

为分析所提架构,我们首先对完整输入提示运行标准全注意力前向传播的Transformer,并记录每层的键值状态。令 C^{(i)} 表示每个层组 G_i 的精确Transformer KV缓存。该矩阵堆叠组内所有层和注意力头的键值状态,并保留 T 个词元位置作为其行。

###### 命题 3.1
令 {φ_r}_{r=1}^T 为沿词元维度的正交归一离散勒让德基,即 φ_r^⊤ φ_{r'}=𝟙_{r=r'}。对于每个 G_i,令 C^{(i)}∈ℝ^{T×D_i}(其中 D_i 为堆叠缓存特征数),且 c_r^{(i)}:=(C^{(i)})^⊤ φ_r∈ℝ^{D_i}(r=1,…,T)。C^{(i)} 的正交分解为 Σ_{r=1}^T φ_r (c_r^{(i)})^⊤。对任意 s>0,有:
$$\displaystyle\left\|C^{(i)}\-\Pi_{N}C^{(i)}\right\|_{F}^{2}\leq\sum_{r=N+1}^{T}\frac{r^{2s}}{(N+1)^{2s}}\left\|c_{r}^{(i)}\right\|_{2}^{2}.$$
这表明误差源于前 N 个模态所省略的勒让德分量。当系数集中于低频时,该尾部较小。增加SSM状态维度 N 可保留更多分量并进一步降低近似误差。证明见附录A(https://arxiv.org/html/2609.17888#A1)。

算法1 通过键、值与输出蒸馏学习SSM
输入:嵌入与查询对 {X,q},参数初始化 Θ={B̄^{(i)},W^{(i)}}_{i=1}^g
要求:分组 {G_i}_{i=1}^g,g 个SSM(各具固定 Ā 及可变权重矩阵 B̄^{(i)})与多层感知机 {f_W^{(i)}}_{i=1}^g,m 层Transformer,n_v,参数 λ₁, λ₂

相似文章

具有自适应退出状态选择的循环状态空间语言模型

arXiv cs.AI

本文探索了使用Mamba和混合Mamba-Transformer骨干网络的循环(递归)状态空间语言模型,表明其在推理任务上优于非循环基线,并在等参数和等FLOPs预训练下保持竞争力,同时自适应退出状态选择改善了中间深度性能。

EndPrompt: 通过终端锚定实现高效长上下文扩展

arXiv cs.CL

EndPrompt 提出了一种方法,仅使用短训练序列即可扩展大语言模型的上下文窗口,通过将终端提示锚定到目标长度的位置索引。该方法在基准测试中取得了优异结果,且计算量远少于全长度微调。

LongAct:利用内在激活模式进行长上下文强化学习

Hugging Face Daily Papers

LongAct 提出了一种显著性引导的稀疏更新策略,通过选择性更新与查询和键向量中高幅值激活相关的权重来改进 LLMs 的长上下文推理能力,在 LongBench v2 上实现了约 8% 的提升。

语言模型推理的选择性状态空间适配与检索

arXiv cs.CL

提出了MaLoRA和MaRA两种适配器系列,在冻结的语言模型中引入选择性状态空间递归,实现token级和上下文级适配,在MuSiQue和2WikiMultihopQA等多跳推理基准上取得了显著提升。

面向切换动态序列的时变深度状态空间模型

arXiv cs.LG

本文提出了一类时变深度状态空间模型,其动态特性通过基函数展开进行学习,从而能够自适应建模切换系统。该方法在合成切换数据和语音去噪任务上均优于时不变模型。