TokenSelect: Efficient Long-Context Inference and Length Extrapolation for LLMs via Dynamic Token-Level KV Cache Selection

发表时间: 2025-11 · EMNLP 2025

原文: https://aclanthology.org/2025.emnlp-main.1079

作者/机构:Wei Wu, Zhuoshi Pan, Kun Fu, Chao Wang, Liyi Chen, Yunchu Bai, Tianfu Wang, Zheng Wang, Hui Xiong / 中国科学技术大学,清华大学,阿里云计算,小红书,香港科技大学(广州),香港科技大学

速读

一句话结论 本文提出了一种免训练的动态 Token 级别 KV Cache 选择方法 TokenSelect,通过多头软投票机制精准筛选关键上下文,在不损失精度的前提下实现了长文本外推,并将注意力计算加速最高 23.84 倍、端到端推理延迟最高加速 2.28 倍。

要解决什么问题 大语言模型在处理长上下文时面临两个卡点:一是预训练长度有限导致直接外推时性能严重下降;二是注意力机制的二次计算复杂度导致推理延迟极高。为了加速,现有方法通常采用稀疏注意力,比如滑动窗口或基于历史分数的 Token 淘汰,或者像 InfLLM 那样进行块级别(Block-level)的 KV Cache 选择。但这些做法卡在了一个关键机制上:真实场景下的注意力稀疏性在 Token 级别是不连续的,且不同注意力头(Attention Head)之间的 Logits 范数差异巨大。块级别的粗粒度选择会漏掉大量离散的关键 Token,而直接按全局分数淘汰又容易被少数数值极大的注意力头主导,导致选出来的上下文并非真正对当前 Query 有用,最终造成长文本信息的不可逆丢失和外推失效。

怎么做的 核心思路是放弃粗粒度的块选择或静态淘汰,改为在每次生成时,针对当前 Query 动态计算历史 KV Cache 中每个 Token 的重要性,并只挑出最关键的少量 Token 参与注意力计算。这样既把计算量降到了常数级,又保留了长程信息。为了绕开“少数头主导”的卡点,方法设计了三个关键部件。首先是多头软投票(Head Soft Vote)机制,它不直接对全局内积排序,而是先计算每个头内部的 QK 内积,通过 Softmax 归一化后再跨头求和,最后选出得分最高的 $k$ 个 Token:

$$ \mathcal{T}_{\mathrm{head-soft-vote}} = \mathrm{TopK} \left( \sum_{h=1}^{H} \sigma \left( \mathbf{Q}^{h} \cdot {\mathbf{K}_{\mathrm{cache}}^{h}}^\top \right) \right) $$

这一设计保证了每个注意力头都能公平地选出自己需要的关键信息。其次,为了消除逐 Token 计算带来的额外开销,作者观察到相邻 Query 之间存在高度的余弦相似性,因此设计了选择缓存(Selection Cache)。在解码阶段,如果当前 Query 与前一个 Query 的相似度大于设定阈值(如 0.9),就直接复用上一步的 Token 选择索引,大幅降低了选择频率。最后,针对底层硬件的显存 I/O 瓶颈,方法结合了 Paged Attention 的分页管理机制,并用 Triton 编写了专用的 Paged Dot Product Kernel。它将选择结果的 I/O 数据量从完整的 KV 向量降维到仅传输索引,并解决了逻辑连续但物理不连续的显存读取问题,从而真正将理论上的复杂度降低转化为了实际的端到端加速。

效果如何 实验在 NVIDIA A100 GPU 上搭建,评测了 Qwen2-7B-Instruct、Llama-3-8B-Instruct 和 Yi-1.5-6B-Chat 三个开源模型,测试数据包含 InfiniteBench、RULER 和 LongBench 三个长文本基准。对比基线涵盖了三大路线:位置编码插值路线(NTK、SelfExtend)、固定稀疏模式路线(StreamingLLM、MInference)以及 KV Cache 动态管理与检索路线(InfLLM、SnapKV、InfiniGen、QUEST、RetrievalAttention)。量化结果显示,在仅保留极小 Token 预算(2K 选中加 512 局部)的设置下,TokenSelect 在 InfiniteBench 上的综合表现大幅超越所有基线,甚至在未经过长文本微调的模型上直接实现了 1M 长度的无损外推。在效率方面,当 KV Cache 长度达到 1M 时,TokenSelect 的注意力计算速度比 FlashInfer 库快了 23.84 倍;在端到端延迟上,单步生成速度比标准全量注意力快 4.70 倍,比同类最优的 InfLLM 快 2.28 倍。作者也坦诚了该方法的局限性:作为一种免训练方案,其绝对性能上限依然受限于基座模型本身的能力(例如 Yi 模型在特定 UUID 检索任务上会因指令遵循能力不足而失效),且在处理 RULER 等极端长文本评测时,整体计算资源消耗依然庞大。

主要贡献

大型语言模型(LLMs)的快速发展激发了现代应用中处理扩展上下文序列的需求。然而,这一进展面临两大挑战:一是由于序列长度超出分布(out-of-distribution)导致的性能下降;二是由于注意力机制的二次计算复杂度导致的推理时间过长。这些问题限制了LLMs在长上下文场景中的应用。

为了解决上述问题,本文的研究目标是提出一种免训练(training-free)的方法,以实现高效且准确的长上下文推理和长度外推。

本文的核心创新点如下:
1. 观察到了注意力机制的非连续稀疏性(non-contiguous attention sparsity),揭示了在Token级别进行KV Cache选择的重要性。
2. 提出了TokenSelect,一种基于动态Token级别的KV Cache选择方法。该方法利用Query和Key的内积来衡量每个注意力头在Token级别的KV Cache关键性,并通过头软投票(head soft vote)机制选择少量关键的KV Cache参与注意力计算,在不牺牲准确性的情况下降低计算量。
3. 基于连续Query之间存在高相似性的观察,设计了选择缓存(Selection Cache),进一步加速了TokenSelect的解码阶段。
4. 实现了高效的分页点积内核(Paged Dot Product Kernel),显著降低了在分页KV Cache管理下的选择开销。

不同稀疏模式下参与注意力计算的token分布(蓝点)。TokenSelect能够更准确地选择关键token(深红色方块)进行注意力计算。
不同稀疏模式下参与注意力计算的token分布(蓝点)。TokenSelect能够更准确地选择关键token(深红色方块)进行注意力计算。

背景知识与关键观察

选择性稀疏注意力问题形式化
LLMs中注意力的高度稀疏性表明,稀疏注意力是解决长上下文推理挑战的有效方案,它可以将参与注意力计算的Token数量保持在恒定规模。由于预定义的稀疏模式会损害性能,本文旨在推理过程中动态选择关键Token。对于长度为$C$的当前输入(在解码阶段$C=1$)和长度为$N$的KV Cache,假设有$H$个大小为$d_h$的注意力头,标准缩放点积注意力(SDPA)的输出$\mathbf{O}$可以表示为:
$\mathbf{O} = \left[ \sigma \left( \frac{\mathbf{Q}^h \cdot \left[ \mathbf{K}_{\mathrm{cache}}^h, \mathbf{K}_{\mathrm{current}}^h \right]^\top}{\sqrt{d}} \right) \cdot \left[ \mathbf{V}_{\mathrm{cache}}^h, \mathbf{V}_{\mathrm{current}}^h \right] \right]_{h=1}^H$
其中$\sigma$表示softmax函数。选择性稀疏注意力的输出$\hat{\mathbf{O}}$则表示为:
$\hat{\mathbf{O}} = \left[ \sigma \left( \frac{\mathbf{Q}^h \cdot \left[ \mathbf{K}_{\mathrm{select}}^h, \mathbf{K}_{\mathrm{current}}^h \right]^\top}{\sqrt{d}} \right) \cdot \left[ \mathbf{V}_{\mathrm{select}}^h, \mathbf{V}_{\mathrm{current}}^h \right] \right]_{h=1}^H$
其中$\mathbf{V}_{\mathrm{select}}^h$是选出的$k$个KV Cache($k \ll N$)。选择过程由选择函数$\mathcal{S}$执行,目标是找到合适的$\mathcal{S}$以最小化$\|\mathbf{O} - \hat{\mathbf{O}}\|_2^2$。现有的工作(如InfLLM、QUEST和MInference)通常采用块级(block-level)选择,这限制了它们的有效性。

注意力在Token级别具有稀疏性、非连续性且各头特征不同
先前的方法(如InfLLM、QUEST和MInference等【引用1, 2, 3】)将KV Cache划分为非重叠的块,并估计块的关键性。这些方法假设关键Token倾向于连续分布。然而,作者观察到这一假设在实践中并不总是成立。如图2a所示,注意力分数在Token级别呈现稀疏分布。这种非连续性导致块级选择出现重大信息遗漏。图2b证明了更细粒度的选择能提高关键Token的召回率,这促使本文采用Token级别的选择。对于Token级选择,直观的方法是直接选择注意力对数(logits)最大的前$k$个Token。但是,图2c揭示了不同注意力头的注意力对数的$L_1$范数存在巨大差异。这会导致选择结果被少数具有极大注意力对数的头所主导,因此需要设计一个更稳健的选择函数来保持各头的独立性。
代币级选择的动机。(a) 注意力分数稀疏性的可视化。(b) 注意力分数和1K代币预算召回的关键代币。(c) 每个注意力头中注意力logits的L1范数。

连续的Query具有相似性
由于注意力的稀疏性是动态的,必须为每个Query执行Token选择,这不可避免地增加了计算开销。作者观察到,连续的Query表现出高度的相似性(如图3a所示)。直观地说,当两个连续的Query高度相似时,它们与Key的内积也会相似,从而导致Token选择结果的大量重叠。作者提出了一个引理:如果两个连续Query的余弦相似度大于阈值$\epsilon$,则它们通过点积选出的前$k$个Key的索引集合是相同的。图3b通过实验证实了这一点,显示Token选择的重叠率随着连续Query相似度的增加而增加。这一关键洞察促使本文在相似的连续Query之间重用选择结果,从而提高计算效率。此外,不同任务间连续Query的相似度分布保持一致,允许在所有场景中应用全局相似度阈值。
连续查询相似性的观察。(a) 连续查询之间的余弦相似度分布。(b) 代币选择重叠率与连续查询相似度的关系。

方法细节

TokenSelect的整体执行流程
TokenSelect的执行流程主要包含三个步骤,如图4所示。首先,通过设计的分页点积内核(Paged Dot Product Kernel)计算每个注意力头的Token级关键性;接着,执行头软投票(head soft vote)机制以获取最终选定的Token索引;最后,通过分页注意力内核(Paged Attention Kernel)仅对选定的稀疏Token执行注意力计算。
TokenSelect的执行流程:1) 通过Paged Dot Product Kernel计算每头关键性;2) 执行头软投票获得选择索引;3) 通过Paged Attention Kernel执行选择性稀疏注意力。

选择函数(Selection Function)的设计演进
最简单的选择函数是通过计算当前$\mathbf{Q}$和历史$\mathbf{K}_{\mathrm{cache}}$的点积来确定Token的关键性,然后选择得分最高的前$k$个Token,即$\mathcal{T}_{\mathrm{topk}} = \mathrm{TopK}((\mathbf{Q} \cdot \mathbf{K}_{\mathrm{cache}}^\top))$。然而,由于不同注意力头之间的注意力对数范数存在巨大差异,这种方法容易产生误差。为了保持各头之间的独立性,一种改进的方法是让每个头各自选出前$k$个最关键的Token,然后通过各头之间的指示函数$\mathbb{I}$进行0/1投票来决定最终选择:$\mathcal{T}_{\mathrm{head-vote}} = \mathrm{TopK}\left(\sum_{h=1}^H \mathbb{I}\left(i \in \mathrm{TopK}\left(\mathbf{Q}^h \cdot {\mathbf{K}_{\mathrm{cache}}^h}^\top\right)\right)\right)$。尽管这种方法性能更好,但它依赖于scatter_add和多次topk操作,导致在GPU上的执行效率极低。此外,0/1投票机制忽略了Token对每个头的相对重要性程度。

头软投票(Head Soft Vote)机制
为了解决上述效率和相对重要性丢失的问题,作者提出了头软投票(head soft vote)方法。具体而言,该方法首先计算每个头(per-head)的关键性得分,然后通过softmax函数$\sigma$将其归一化,最后将所有头的归一化得分相加,并提取Top-K索引:$\mathcal{T}_{\mathrm{head-soft-vote}} = \mathrm{TopK}\left(\sum_{h=1}^H \sigma\left(\mathbf{Q}^h \cdot {\mathbf{K}_{\mathrm{cache}}^h}^\top\right)\right)$。这种设计不仅平衡了各头的贡献,避免了被极端值主导,而且在GPU上具有更高的执行效率。

优化预填充(Prefill)阶段的选择频率
虽然选择函数将注意力的复杂度从$O(N^2)$降低到了$O(k^2)$,但选择函数本身的执行时间仍会影响推理延迟。在预填充阶段,输入的Query矩阵为$\mathbf{Q}_{\mathrm{prefill}} \in \mathbb{R}^{n_{\mathrm{in}} \times d}$。在长上下文场景中,用户输入序列的Token数量$n_{\mathrm{in}}$可能高达100万(1M),对每个Query Token逐一执行选择是不切实际的。考虑到连续Query的相似性,TokenSelect采用了分块(chunk-wise)Token选择策略。具体操作是将一个大小为$c$的查询块$\mathbf{Q}_C \in \mathbb{R}^{c \times d}$内的所有Query取平均值,即$\frac{1}{c} \sum_{i=1}^c (\mathbf{Q}_C)_i$,然后将该平均向量输入到选择函数中。这种方法有助于保持预填充阶段的计算密集型特征,防止其退化为受限于内存带宽(memory bound)的操作。

优化解码(Decode)阶段的选择频率(Selection Cache)
在解码阶段,由于LLMs的自回归生成特性,需要为每一个新生成的$\mathbf{Q}_{\mathrm{decode}}$频繁执行Token选择,且无法像预填充阶段那样进行分块处理。为了降低解码阶段的Token选择频率,作者提出了选择缓存(Selection Cache)。当连续的Query高度相似(余弦相似度大于设定的阈值$\theta$)时,将命中该缓存,系统会直接加载并复用前一个Query的Token选择结果,而跳过当前步的选择计算。选择缓存能够在保证模型性能的同时,有效降低解码延迟。

选择性稀疏注意力的效率瓶颈与高效实现
为了使TokenSelect适用于实际的生产环境,高效的底层实现是必不可少的。作者首先分析了代表性的块级选择性稀疏注意力方法InfLLM【引用2】的时间开销。如图5所示,尽管降低了理论计算复杂度,但实际运行时间严重依赖于实现方式。由于需要依赖历史注意力分数,许多方法(如$\mathrm{H_2O}$、TOVA、SnapKV、InfLLM【引用2, 4, 5】)无法与Flash Attention等高效内核兼容,导致在实际服务中不可用。进一步分析InfLLM的Flash Attention兼容版本发现,点积计算本身并不是主要瓶颈;相反,在更新块和拼接KV Cache期间,在GPU显存(HBM)中索引和合并选定的KV Cache Token会产生极大的I/O开销,这加剧了LLM推理的内存受限问题。
不同注意力实现下单个块预填充步骤的时间分解(块大小:512,KV Cache长度:128K,参与注意力的token:4K)。

分页点积内核(Paged Dot Product Kernel)设计
基于上述I/O瓶颈分析,作者提出分页注意力(Paged Attention)是实现选择性稀疏注意力更合适的底层架构。通过使用分页KV Cache管理(在TokenSelect中页面大小设置为1),可以将选择结果的I/O数据量从所有选定KV Cache的规模$O(2kd)$大幅减少到仅传输其索引的规模$O(k)$。然而,图5中的第(4)项揭示了在分页KV Cache管理下的另一个新瓶颈:由于逻辑上连续的KV Cache在物理HBM中并不完全连续,在执行选择操作之前需要将其转换为连续格式。为解决这一问题,作者利用Triton编写了定制的分页点积内核(Paged Dot Product Kernel),该内核允许直接在非连续的分页显存上高效执行点积运算,从而显著提升了TokenSelect的整体执行效率。

实验环境

实验结果

结论

本文提出了TokenSelect,一种用于高效长上下文推理和长度外推的免训练方法。TokenSelect通过新颖的Token级选择性稀疏注意力机制,成功解决了LLMs在处理长文本时面临的两个主要挑战:预训练带来的上下文长度限制以及注意力机制的二次计算复杂度。实验结果表明,TokenSelect在注意力计算上可实现高达23.84倍的加速,在端到端推理延迟上实现高达2.28倍的加速,同时在多个长上下文基准测试中展现出卓越的性能。

补充细节(局限性)

尽管TokenSelect取得了显著成果,但仍存在一些局限性,为未来工作提供了方向。首先,其免训练设计虽然是一大优势,但也是一把双刃剑,因为其绝对性能本质上与底层LLMs的质量绑定。例如在实验中发现Yi-1.5-6B-Chat错误地识别了UUID字符串,这表明某些问题仍需要通过模型微调来解决。其次,虽然TokenSelect在长上下文推理中达到了最先进的性能,但LLM社区最近的长文本后训练技术也展现了令人印象深刻的效果;TokenSelect与这些方法是正交的,可以在推理期间结合使用,以轻微的性能下降换取显著的效率提升。最后,尽管该方法大幅提升了效率,长上下文推理本质上仍然是资源密集型的。例如,即使使用8B参数模型,运行复杂的RULER基准测试仍需要大约8张A100 GPU运行近一天时间,更大模型的计算成本将更加高昂。需要社区在模型设计、算法开发和基础设施优化方面共同进步,以进一步缓解这些计算挑战。

附录细节

算法1:选择缓存算法(Selection Cache Algorithm)
选择缓存算法旨在通过在相似Query之间共享选择结果来降低选择频率。算法输入包括当前查询向量$\mathbf{Q}$、需选择的Token数$k$、缓存的查询向量$\mathbf{C}_Q$、缓存的索引$\mathbf{C}_\mathcal{T}$、余弦相似度阈值$\theta$等。执行时,如果当前是第一个查询标志为真,或者当前Query与缓存Query的余弦相似度低于阈值($\cos(\mathbf{Q}, \mathbf{C}_Q) < \theta$),则调用选择函数$\mathcal{S}(\mathbf{Q}, k)$计算新的Top-$k$索引$\mathcal{T}$,并将结果更新至缓存$\mathbf{C}_\mathcal{T}$和$\mathbf{C}_Q$中;反之,如果相似度大于等于阈值,则直接复用缓存中的索引$\mathcal{T} = \mathbf{C}_\mathcal{T}$。最后返回索引$\mathcal{T}$。

算法2:分页点积内核(Paged Dot Product Kernel)
该算法详细描述了如何在分页KV缓存管理下高效执行Token级单头关键性估计,以显著减少HBM和SRAM之间的I/O通信。算法并行遍历每个注意力头$h$,将当前头的查询向量$q$加载到SRAM中。接着,并行遍历相关的Token块(CUDA block大小为$B$),将该块内的Token索引以及对应的键向量$k$从HBM加载到SRAM中。在SRAM内部直接计算点积$s = \langle q, k \rangle$,最后将计算得到的关键性分数$\mathbf{S}$写回HBM。

扩展至100万以上上下文长度的可扩展性
为了进一步探索TokenSelect在极端长上下文场景中的性能,作者在InfiniteBench的基础上设计了不同文本长度的扩展基准测试。如图9所示,TokenSelect能够在高达200万(2M)Token的上下文中,仅使用极小的Token预算就能准确捕获关键信息,突显了其在更广泛应用场景中的潜力。
使用Qwen2-7B-Instruct在扩展的R.PK和R.KV上的性能比较。

扩展至720亿参数模型的可扩展性
为了证明该方法在更大模型上的可扩展性,作者使用Qwen2-72B-Instruct模型进行了额外实验(配置张量并行度为4)。结果表明,TokenSelect在准确率和延迟方面均优于NTK-Aware Scaled RoPE基线,证明了该方法能够有效扩展至大规模语言模型。

引理1的正式声明与证明
附录C提供了关于“余弦相似度阈值下Top-k键选择不变性”的正式证明。假设存在两个查询向量$\mathbf{q}_1, \mathbf{q}_2$和一组键向量$\mathbf{k}_i$,以及基于点积相似度的Top-k选择函数$\mathcal{T}(\mathbf{q})$。证明过程通过引入正交分解$\hat{\mathbf{q}}_2 = \hat{\mathbf{q}}_1 \cos\theta + \hat{\mathbf{p}}_1 \sin\theta$,推导出了当$\cos(\mathbf{q}_1, \mathbf{q}_2) > \epsilon$时,$\mathcal{T}(\mathbf{q}_1) = \mathcal{T}(\mathbf{q}_2)$成立的充分条件,为选择缓存(Selection Cache)的设计提供了理论依据。

与基于Token驱逐(Token Eviction)方法的比较
附录E详细比较了TokenSelect与以$\mathrm{H_2O}$【引用4】为代表的Token驱逐方法。尽管两者都采用Token级关键性估计,但$\mathrm{H_2O}$属于独立于查询(query-independent)的KV缓存选择方法,存在三个主要缺点:1)缺乏动态性,其重要性评分依赖于历史查询,可能导致当前查询急需的历史KV对已被提前丢弃;2)无法扩展序列长度,因为它依赖模型原始注意力机制;3)实现效率低下,基于注意力分数的评估使其不兼容FlashAttention等高效内核。相比之下,TokenSelect采用基于当前查询的动态选择策略,能够轻松将有效上下文扩展至1M以上,且完全透明兼容分页注意力、张量并行等大规模推理加速基础设施。

预填充(Prefill)延迟比较
附录H对比了TokenSelect与广泛应用于实际长上下文推理的MInference【引用1】的端到端预填充延迟。在使用Llama-3-8B和单张A100的测试中,TokenSelect在较短输入Token长度(如1K, 10K)下表现出显著的延迟优势(0.092s vs 3.017s),并在输入长度增加时保持了与MInference相当的效率。更为突出的是,当序列长度达到200K和300K导致MInference发生内存溢出(OOM)时,TokenSelect依然能够顺利完成预填充计算。

参考文献引用汇总