SepLLM: Accelerate Large Language Models by Compressing One Segment into One Separator

发表时间: 2025-07 · arXiv:2412.12094 (ICML 2025)

原文: https://arxiv.org/abs/2412.12094

作者/机构:Guoxuan Chen 1 2, Han Shi 1, Jiawei Li 1, Yihang Gao 2, Xiaozhe Ren 1, Yimeng Chen 3, Xin Jiang 1, Zhenguo Li 1, Weiyang Liu 4, Chao Huang 2

速读

一句话结论
提出了一种名为 SepLLM 的稀疏注意力框架,通过将长文本片段的语义信息强制压缩到标点符号等分隔符中,在仅保留初始、局部和分隔符 Token 的情况下,实现了大模型训练与推理的双重加速,在 Llama-3-8B 上减少了超过 50% 的 KV Cache 且复杂推理性能几乎无损。

要解决什么问题
标准 Transformer 的自注意力机制计算复杂度和显存占用随序列长度呈平方级增长,这导致长上下文场景下 KV Cache 显存爆炸和推理延迟急剧升高。现有的优化路线存在明显的机制卡点:第一类线性注意力机制会彻底改变模型架构,无法直接复用当前强大的开源预训练权重;第二类基于 KV Cache 压缩的免训练方法(如保留局部和初始 Token 的 StreamingLLM)由于在生成过程中粗暴地丢弃了大量中间 Token,导致模型在长文本理解和多步逻辑推理上的精度严重掉点。此外,像 H2O 或 SnapKV 这种基于注意力分数动态淘汰 Token 的策略,绝大多数难以无缝接入模型的训练阶段,造成了训练时全注意力与推理时稀疏注意力的行为分布不一致,限制了模型长文本能力的理论上限。

怎么做的
核心思路是利用自然语言的天然分布特性,将长文本片段的语义信息压缩到用于断句的分隔符(如逗号、句号、换行符等)中。作者观察到,在标准大模型的注意力图里,这些看似无意义的分隔符占据了极高的注意力权重。基于此,SepLLM 设计了一种数据依赖的稀疏注意力机制,强制当前 Token 在计算注意力时,只能看到三种历史信息:固定数量的初始 Token(用于稳定注意力的 Attention Sinks)、距离当前位置最近的 $n$ 个局部相邻 Token(捕捉局部平滑语义),以及当前位置之前的所有分隔符 Token(作为各个历史片段的语义压缩包)。

这种设计之所以能绕开上述卡点,是因为它既不像全注意力那样保留所有冗余 Token 导致显存溢出,又不像 StreamingLLM 那样彻底丢失中间片段的全局信息。其注意力掩码矩阵 $\mathbf{M}$ 的更新机制定义如下:

$$\begin{aligned} \Lambda_{i,j} = \begin{cases} \mathbf{Q}_i^\top \mathbf{K}_j / \sqrt{d_k}, & \text{if } \mathbf{M}_{i,j} = 1 \\ -\infty, & \text{if } \mathbf{M}_{i,j} = 0 \end{cases} \end{aligned}$$

$$ \mathbf{O} = \text{Softmax}(\Lambda) \cdot \mathbf{V} $$
其中,当第 $j$ 个 Token 属于初始、局部或分隔符时,掩码 $\mathbf{M}_{i,j} = 1$,否则为 $0$。

为了支持无限长度的流式输入,SepLLM 设计了四个动态缓存区:初始缓存、局部窗口缓存、历史窗口缓存和分隔符缓存。当总缓存达到容量上限时,系统会触发压缩机制,将历史窗口中的分隔符提取到分隔符缓存中,并直接丢弃其余普通 Token。更关键的是,该方法不仅用于推理,作者还基于 FlexAttention 开发了底层算子 Sep-Attention,使其能直接嵌入到从头训练或微调阶段,强制模型在训练时就学会将片段信息汇聚到分隔符上,彻底消除了训练与推理的鸿沟。

效果如何
实验覆盖了免训练、从头训练和训练后微调三种设置。模型规模涵盖 Pythia-160m 到 Falcon-40B,重点测试了 Llama-3-8B。对比基线包括代表全注意力路线的 Vanilla 模型,代表保留局部与初始 Token 路线的 StreamingLLM,以及代表动态 KV Cache 压缩路线的 H2O、SnapKV 和 PyramidKV。

在免训练设置下,基于 Llama-3-8B 测试 GSM8K-CoT 数学推理任务,SepLLM 在仅使用 47.36% 运行时 KV Cache 的情况下,取得了与全注意力基线几乎一致的准确率(77.18% 对比 77.79%),而同等缓存限制下的 StreamingLLM 准确率暴跌至 69.67%。在 MMLU 知识推理基准上,SepLLM 同样在仅用 44.61% 缓存时保持了 64.68% 的得分,显著优于 StreamingLLM。在从头训练设置下(使用 300B tokens 训练 Pythia-160m),SepLLM 相比标准 Transformer 减少了约 30% 的计算量和 26% 的训练时间,且在相同计算开销下达到了更低的 Loss。在 PG19 数据集的流式长文本测试中,SepLLM 能够稳定处理高达 400 万 Token 的超长序列,且困惑度始终低于 StreamingLLM。在“大海捞针”测试中,SepLLM 也能成功检索到目标信息,证明分隔符确实起到了信息压缩的作用。

该方法的局限在于,其性能对分隔符的定义有一定依赖,消融实验表明,如果仅保留句号和问号而剔除逗号等其他符号,模型的推理能力会出现明显下降;此外,为了达到最优的下游任务表现,通常需要采用混合层设计,即把模型的第一层和最后一层保持为全注意力机制。

主要贡献

基于Transformer的模型在各类任务中表现出色,但依赖于下一个token预测的标准Transformer面临着巨大的计算挑战,尤其是在扩展到更大模型和更长上下文时。这些计算效率低下的问题主要源于自注意力模块对输入token数量的二次复杂度,这显著影响了推理速度和训练时间。现有的线性注意力方法会大幅改变架构,导致无法直接利用强大的预训练模型;而免训练的KV Cache优化方法(如StreamingLLM)在训练阶段适应性差,会导致训练和推理性能之间存在差异。

为了更好地理解大语言模型(LLMs)的内在机制,作者分析了不同样本的注意力模式。研究发现,LLMs在进行信息检索时,往往不会将注意力集中在具有实际语义的token(如名词和动词)上,而是不成比例地高度关注看似“无意义”的分隔符token(如“.”或“\n”)。这一观察表明,片段信息被压缩并嵌入到了这些分隔符token中,从而实现了高效的信息检索,而无需直接从内容token中提取。

基于这一洞察,本文提出了SepLLM,这是一个即插即用的框架,通过压缩这些片段并消除冗余token来加速推理。同时,作者还实现了用于训练加速的高效算子。主要贡献总结如下:
* 通过可视化token级别的注意力分数,揭示了初始、相邻和分隔符token始终获得较高的注意力权重,并据此提出了SepLLM框架。
* 通过在充分训练的LLMs上进行针对性的掩码实验,证明了分隔符token包含关键信息且对模型性能至关重要。
* 在免训练、从头训练和后训练设置下进行了全面实验,证明了SepLLM在各种任务、数据集和主干模型上的有效性。例如,在Llama-3-8B主干上,SepLLM在GSM8K-CoT基准测试中减少了50%以上的KV Cache,同时保持了相当的性能;在流式设置中,能有效处理长达400万以上的token序列。
* 开源了代码库,支持使用Sep-Attention模块进行高效的多节点分布式训练,并兼容多种融合算子(如fused rope、fused layer norm等)。

图1 香草Transformer与提出的SepLLM之间的损失比较。在不同的计算成本和不同的训练时间下,SepLLM始终实现了更低的损失。
图1 香草Transformer与提出的SepLLM之间的损失比较。在不同的计算成本和不同的训练时间下,SepLLM始终实现了更低的损失。
图2 输入“Natalia sold clips to 48 of her friends in April, and then she sold half as many clips in May. ...”时,不同层的注意力分数可视化。注意,像“,”和“.”这样的分隔符token贡献了大量的注意力。
图2 输入“Natalia sold clips to 48 of her friends in April, and then she sold half as many clips in May. ...”时,不同层的注意力分数可视化。注意,像“,”和“.”这样的分隔符token贡献了大量的注意力。

背景知识与设计原则

KV Cache压缩的局限性。近期的研究致力于克服LLMs在处理超长上下文输入时的限制。FastGen提出了自适应KV Cache管理方法;SnapKV利用注意力分数选择并聚类重要位置;$H_2O$ 实施了动态token保留策略;StreamingLLM通过保留attention sinks和局部token,使LLMs能够处理无限长度序列而无需微调;PyramidInfer和PyramidKV则修改了不同层的KV Cache容量。然而,这些类别中的绝大多数工作都无法应用于训练阶段。

稀疏注意力的特性。稀疏注意力通过将注意力限制在预定义的模式(如局部窗口或固定步长的块模式)来创建稀疏注意力矩阵。Longformer结合了空洞局部窗口注意力与任务特定的全局注意力;BigBird提出了使用全局token、局部滑动窗口和随机注意力的线性复杂度替代方案;SparseBERT提出了一种可微的注意力掩码算法来端到端地学习注意力掩码。需要注意的是,大多数关于稀疏注意力的工作都使用固定的掩码,并且是建立在BERT家族之上的。相比之下,本文提出的SepLLM主要建立在GPT系列之上,且其注意力掩码是依赖于数据的(data-dependent)。

方法细节

基础架构设计

注意力掩码的重构设计。从图2中可以观察到,在给定的输入上下文中,看似“无意义”的分隔符token比具有实际语义的token获得了更高的注意力分数。因此,作者提出了一种新颖的Transformer架构,在Transformer的某一层(即自注意力层)中,输入中的每个token只能看到前一个Transformer层输出的、位于当前token之前的一部分(而非全部)token的隐藏状态。这个token子集包括一定数量的初始词(例如attention sinks)、当前token之前的所有分隔符token,以及距离当前token最近的 $\pmb{n}$ 个token。

初始Token的保留机制。当使用滑动窗口机制进行生成时,移除KV Cache中对应于初始token的键值对(KV)会导致生成token的困惑度显著增加,这是StreamingLLM 【索引1,Efficient Streaming Language Models with Attention Sinks+2024+ICLR+https://arxiv.org/abs/2309.17453】中提到的一种现象。最初的几个token也被称 为attention sinks。作者保留了这一设置,并在随后的实验中进一步验证了初始token的作用。通常情况下,会保留 $\pmb{a}$ 个初始token。

分隔符Token的信息压缩假设。从图2中可以观察到,在给定的输入上下文中,分割序列的看似“无意义”的分隔符token(如逗号、句号、感叹号、分号等)比具有语义的token(如名词或动词)获得了更高的注意力分数。因此,作者假设这些分隔符自然地压缩了由它们分割的文本段落的信息,这样当Transformer生成新token时,它只需要引用这些分隔符中包含的信息,就能提取与那些文本段落相关的信息。因此,在免训练场景中,作者采用了这种策略,并在许多任务上取得了与基于全注意力的原始模型相似的结果。此外,为了强化使用分隔符压缩其各自段落内信息的效果,作者采用了从头训练和后训练的方法,迫使模型在训练期间限制当前token访问远处前置文本的所有信息,即在每个段落中,当前token只能看到代表其段落的分隔符(其他token被掩码,见图3)。以这种方式训练后,段落内的信息被迫浓缩到分隔符中,导致Transformer预测下一个词的概率分布与具有全注意力的原始Transformer非常相似。

图3 SepLLM的整体范式。左侧展示了在给定输入“ABC, DE. FG \n”时,训练或预填充阶段的注意力掩码。右侧展示了生成阶段的KV Cache管理。
图3 SepLLM的整体范式。左侧展示了在给定输入“ABC, DE. FG \n”时,训练或预填充阶段的注意力掩码。右侧展示了生成阶段的KV Cache管理。

相邻Token的局部依赖捕获。语言任务通常表现出很强的局部依赖性和交互性,因为相邻的token通常构成连贯的短语,或者具有需要被捕获的依赖关系。相邻token通常有助于形成局部平滑连贯的上下文,使模型能够生成在直接上下文中合理的句子。相邻token(也被称为局部注意力或滑动窗口注意力)在各种高效Transformer(如StreamingLLM和$H_2O$)中都被考虑过,作者也采用了这种方法,将距离当前token最近的前置token数量记为 $\pmb{n}$。

整体处理流程

训练与预填充阶段的稀疏矩阵乘法。在SepLLM架构的训练/预填充阶段,不需要将输入上下文中所有token对应的查询向量与所有的键向量相乘。只需要将图3中掩码矩阵中高亮元素对应的query-key对的向量相乘即可。该公式可以表示如下:

$$\begin{aligned} \begin{array} { c } { \displaystyle \mathbf { A } = \mathrm { Softmax } \left( \Lambda \right) , \Lambda { = } \frac { \mathbf { Mul } \left( \mathbf { Q } , \mathbf { K } ^ { \top } \big | \mathbf { M } \right) } { \sqrt { d _ { k } } } \ } \\ { \mathbf { O } = \mathbf { A } \cdot \mathbf { V } } \end{array} \end{aligned}$$

其中 $\mathbf{Q} \in \mathbb{R}^{m \times d_k}$,$\mathbf{K} \in \mathbb{R}^{m \times d_k}$ 是一层注意力层中的查询和键矩阵,其中每个行向量 $\mathbf{Q}_i, \mathbf{K}_j$ 对应于序列长度为 $m$ 的输入上下文中第 $i$ 个token的查询和第 $j$ 个token的键。$d_k$ 表示键和查询向量的维度。

图4 为流式应用量身定制的SepLLM整体框架。KV对存储在四个缓存块中。图示展示了提出的SepLLM的流式处理过程。每一行代表一次迭代,四个缓存块动态更新(显示为四列),并在每次迭代中更新(显示在单行中)。一旦运行时使用量 $Size_{run}$ 达到最大容量 $c$,SepLLM将过去窗口缓存中分隔符token的KV Cache移动到分隔符缓存中,并丢弃其他KV Cache。
图4 为流式应用量身定制的SepLLM整体框架。KV对存储在四个缓存块中。图示展示了提出的SepLLM的流式处理过程。每一行代表一次迭代,四个缓存块动态更新(显示为四列),并在每次迭代中更新(显示在单行中)。一旦运行时使用量 $Size_{run}$ 达到最大容量 $c$,SepLLM将过去窗口缓存中分隔符token的KV Cache移动到分隔符缓存中,并丢弃其他KV Cache。

掩码矩阵与输出计算。$\boldsymbol{\Lambda}, \mathbf{A} \in \mathbb{R}^{m \times m}$ 分别是原始和最终的注意力图。$\mathbf{V} \in \mathbb{R}^{m \times d_v}$ 是维度为 $d_v$ 的值矩阵,$\mathbf{O} \in \mathbb{R}^{m \times d_v}$ 表示当前注意力层的输出。$\mathrm{Mul}(\cdot)$ 表示稀疏矩阵乘法函数,该函数可以通过SampleAttention等方法进行优化,作者也实现了名为Sep-Attention的自有模块来加速此过程。$\mathbf{M} \in \mathbb{B}^{m \times m}$ 是一个二值掩码矩阵,作为 $\mathrm{Mul}(\cdot)$ 的参数:

$$\begin{aligned} \begin{array} { r l r } & { } & { \Lambda _ { i , j } = \left\{ \begin{array} { l l } { { \bf Q } _ { i } ^ { \top } { \bf K } _ { j } / \sqrt { d _ { k } } , } & { \mathrm { if } \ { \bf M } _ { i , j } = 1 } \\ { \quad \ - \infty , } & { \mathrm { if } \ { \bf M } _ { i , j } = 0 } \end{array} \right. . } \end{array} \end{aligned}$$

其中 $\Lambda_{i, j}, \mathbf{A}_{i, j}, \mathbf{M}_{i, j}$ 分别是矩阵 $\Lambda, \mathbf{A}, \mathbf{M}$ 第 $i$ 行第 $j$ 列的元素。因为如果 $\Lambda_{i, j} = -\infty$ 则 $\mathbf{A}_{i, j} = 0$,所以非初始、非分隔符和非相邻的token将被 $\mathbf{A} \cdot \mathbf{V}$ 掩码。这种策略适用于多头注意力的所有头。

生成阶段的缓存保留策略。在生成阶段,这种基础设计的KV Cache管理也非常直观。如图3右侧所示,在生成新token时,仅保留初始、分隔符和相邻token的KV Cache。因此,SepLLM中的KV Cache要小得多,需要的内存也更少。理想情况下,基于SepLLM,生成下一个词的困惑度与具有全注意力的原始Transformer相当。

定制的流式设计

流式场景中的无尽累积问题。在真实世界场景中,存在许多流式应用(如多轮对话),这些应用预期会有很长的交互。因此,期望SepLLM能够处理无限输入而不会显著牺牲效率和性能。正如基础设计中所讨论的,SepLLM可以通过仅保留分隔符、相邻和初始token的KV来节省大量的KV Cache。然而,随着输入token数量的增加,KV Cache中的分隔符数量也会无休止地累积,这对于流式设置是不可行的。因此,针对流式场景提出了定制的流式设计。

四区块缓存系统框架。图4展示了流式应用中SepLLM的处理架构。系统同时维护四个专门的缓存块:初始缓存(Initial Cache)、分隔符缓存(Separator Cache)、过去窗口缓存(Past Window Cache)和局部窗口缓存(Local Window Cache)。具体来说,初始缓存捕获attention sinks。局部窗口和过去窗口缓存存储连续token的KV,过去窗口缓存作为局部窗口缓存的溢出缓冲区。分隔符缓存保留包含压缩片段信息的分隔符的KV。

缓存容量参数定义。为了描述缓存管理策略,将四个缓存的运行时使用量分别表示为 $Size_{\mathrm{init}}$、$Size_{\mathrm{sep}}$、$Size_{\mathrm{past\_w}}$ 和 $Size_{\mathrm{local\_w}}$。所有KV Cache的运行时使用量定义为 $Size_{\mathrm{run}} := Size_{\mathrm{init}} + Size_{\mathrm{sep}} + Size_{\mathrm{past\_w}} + Size_{\mathrm{local\_w}}$,其满足 $Size_{\mathrm{run}} \leq \pmb{c}$。连续相邻token的数量定义为 $\pmb{n} := Size_{\mathrm{past\_w}} + Size_{\mathrm{local\_w}}$。值得注意的是,在流式设置中,$\pmb{n}$ 是输入序列长度 $m$ 的函数,而不是固定的超参数。该缓存系统的预设超参数如下(注意:$\pmb{a} + \pmb{s} + \pmb{w} < \pmb{c}$):
* $\pmb{c}$:整个KV Cache的最大容量。
* $\pmb{a}$:初始缓存的最大容量。
* $\pmb{s}$:分隔符缓存的最大容量。
* $\pmb{w}$:局部窗口缓存的最大容量。值得注意的是,当运行时KV Cache使用量 $Size_{\mathrm{run}}$ 首次达到 $\pmb{c}$ 后,$\pmb{w}$ 也是 $\pmb{n}$ 的最小值。

流式生成中的动态缓存更新过程。在流式序列生成期间,SepLLM将首先填充初始缓存,然后填充局部窗口缓存。在 $Size_{\mathrm{local\_w}}$ 达到 $\pmb{w}$ 后,后续的token被引导至过去窗口缓存。当 $Size_{\mathrm{run}}$ 达到 $\pmb{c}$ 时(图4中的迭代1),触发压缩机制,此时过去窗口缓存中的分隔符token被移动到分隔符缓存中,其他token被丢弃。当总输入长度达到某个 $m_0$ 使得分隔符缓存达到其容量 $\pmb{s}$ 时,$\pmb{n}$ 进入周期模式。具体而言,对于 $m > m_0$,$\pmb{n}$ 遵循以 $\pmb{w}$ 和 $\pmb{c} - \pmb{a} - \pmb{s}$ 为界的周期性线性函数。无限长序列生成的平均运行时KV Cache使用量满足:

$$\begin{aligned} \begin{array} { r } { \underset { m \to \infty } { \operatorname* { lim } } \overline { { Size _ { \mathrm { run } } } } = \underset { m \to \infty } { \operatorname* { lim } } \bar { \pmb { n } } _ { m } + \pmb { a } + \pmb { s } } \\ { = \frac { \pmb { w } + \pmb { c } + \pmb { a } + \pmb { s } } { 2 } < \pmb { c } . } \end{array} \end{aligned}$$

位置编码偏移策略。流式设置的位置编码策略与最先进的StreamingLLM相同,专门为无限长输入设计,即关注缓存内的位置而不是原始文本中的位置。

实验环境

  • 模型架构与参数:采用Pythia(Biderman et al., 2023)和Llama-3(Dubey et al., 2024)两个模型系列。具体而言,Pythia-160m-deduped用于从头训练任务;Pythia-1.4B-deduped用于后训练设置;Llama-3-8B-Instruct用于免训练和流式任务。
  • 数据集名称、规模及用途

    • 在从头训练和后训练任务中,使用去重后的Pile数据集(包含约207B个token)。从头训练设置为1.5个epoch(全局批次大小为1024,共143000步),总计使用了约300B个token。
    • 评估基准包括:GSM8K-CoT(测试数学推理和逻辑步骤,8-shot)、MMLU(测试多学科知识,5-shot)、ARC、LAMBADA、LogiQA、PIQA、SciQA,以及用于长文本流式测试的PG19和WikiText。
  • 分隔符配置:所有评估中使用的分隔符token包括:[".", ",", "?", "!", ";", ":", " ", "\t", "\n"]。

实验结果

1. 免训练(Training-Free)实验
* 实验内容:基于Llama-3-8B-Instruct模型,在GSM8K-CoT和MMLU基准上评估SepLLM和StreamingLLM(StrmLLM)。
* 实验结果:SepLLM (n=256) 在GSM8K-CoT(77.18%)和MMLU(64.68%)上取得了与全注意力Llama-3几乎相同的性能,但仅使用了全注意力模型约47%的KV Cache。相比之下,移除了分隔符KV的StrmLLM (n=256) 性能显著下降。即使将StrmLLM的窗口大小增加到n=380以使其KV Cache使用量与SepLLM相当,其性能依然低于SepLLM。
* 分析结论:分隔符的KV确实封装了其各自段落内包含的信息,移除它们会严重影响Transformer的理解和推理能力。

2. 从头训练(Training from Scratch)实验
* 实验内容:在Pile数据集上从头训练Pythia-160m-deduped模型,对比全注意力(Vanilla)、StrmLLM和SepLLM的不同配置。
* 实验结果
* 相邻Token的益处:增加相邻Token数量(从n=64到n=128),训练损失下降更快(图5),下游任务性能更强(表2)。
* 混合层的益处:将第一层和最后一层改为全注意力(SepLLM (n=64, H/T)),能进一步优化训练过程和下游任务表现。
* 分隔符的作用:StrmLLM (n=64) 训练损失下降明显变慢,且下游任务性能恶化。
* 计算效率:SepLLM可显著减少约30%的FLOPs(表3)。在相同的FLOPs下,SepLLM的损失低于Vanilla(图5b)。

  • 分析结论:SepLLM架构在训练期间提取有用信息的能力至少与Vanilla相当,并且能显著加速训练过程。
图5 从头训练的训练损失曲线。图5(b)显示了不同方法在相同FLOPs下相对于Vanilla的损失比率。
图5 从头训练的训练损失曲线。图5(b)显示了不同方法在相同FLOPs下相对于Vanilla的损失比率。

3. 后训练(Post-Training)实验
* 实验内容:使用Pythia-1.4B-deduped的93000步检查点进行后训练。并在PG19测试集上(生成20K和64K token)评估端到端推理时间、困惑度和KV Cache使用量。
* 实验结果:图6显示,增加n和适当提高学习率都有助于损失下降。在PG19上(表4和表5),在相同的最大KV Cache容量下,SepLLM预测下一个token的平均困惑度始终低于StreamingLLM,且推理时间更短。
* 分析结论:SepLLM可以通过后训练迅速将全注意力模型转换为适应SepLLM架构嵌入分布的模型。

图6 后训练设置的训练损失曲线。
图6 后训练设置的训练损失曲线。

4. 流式应用与消融实验(Streaming Applications & Ablation Study)
* 实验内容:在WikiText上针对长输入应用进行超参数(s, c, w)消融;测试初始token和位置编码偏移的影响。
* 实验结果
* 增加分隔符缓存容量(s)、总容量(c)和窗口大小(w)均能降低长文本推理的困惑度。
* 移除初始token会严重影响困惑度(表8)。
* 移除位置编码偏移后,StreamingLLM的困惑度飙升至400以上,而SepLLM仅上升至200左右。

  • 分析结论:SepLLM对位置编码偏移的依赖程度较低,进一步凸显了分隔符在预测token时的稳定性作用。

5. 与更多基线及变体的比较
* 实验内容:在MMLU上将SepLLM与朴素的head-wise滑动窗口基线进行比较。同时提出了一个固定间隔变体(FixLLM),每隔固定数量的token保留一个注意力。
* 实验结果:朴素基线表现极差(表9)。FixLLM在数学逻辑和知识推理能力上均显著落后于SepLLM。
* 分析结论:SepLLM中分隔符对片段信息的压缩和总结能力是固定间隔token无法替代的。

结论

本文致力于通过高效的神经架构修改来解决LLMs在处理长输入时的计算和存储挑战。通过注意力图的可视化,发现特定的分隔符token始终贡献很高的注意力分数。受此启发,提出了SepLLM,这是一种新的语言建模视角和稀疏注意力机制,将注意力计算集中在初始、相邻和分隔符Token上。为了实现挂钟时间的加速,还实现了硬件高效的算子。免训练研究表明,这些分隔符有效地压缩了片段信息,实现了高效的信息检索。与以往的免训练方法不同,SepLLM可以被整合到训练阶段(如从头训练或后训练),从而减少了训练和推理之间的差异。在各种设置下进行的大量实验证明了SepLLM的实际有效性。

附录与补充细节

补充细节:通用近似定理(Universal Approximation)

理论分析证明了基于编码器的SepLLM具备通用近似能力。定理5.1指出,给定 $p > 1$ 和 $n > 2$,对于任何连续函数 $f \in \mathcal{F}$ 和 $\epsilon > 0$,都存在一个SepLLM $g$,使得它们之间的距离 $d_p(f, g) < \epsilon$。证明过程(见附录J和K)通过将输入空间网格化,并利用SepLLM的注意力层和前馈层构造量化映射、上下文映射和值映射,最终证明了SepLLM可以以任意精度逼近任何序列到序列的连续函数。

补充细节:原生稀疏注意力与语言建模

SepLLM本质上是基于自然语言的自然语义分布来建模稀疏性的。在预训练阶段,SepLLM有意将片段信息压缩到用于划分该片段的分隔符中。这种方法与自然语言的语义分布紧密对齐,因为分隔符本身提供了对当前片段的划分和总结。被分离出来的片段在语义上本质是连贯的,形成了自包含的语义单元。因此,可以将SepLLM视为自然语言本身固有的一种原生稀疏注意力机制。

附录A:注意力分数可视化

使用了Llama-3-8B-Instruct模型,输入包含数学推理的句子。图12、13、14展示了不同层和头的注意力图,证实了LLM在处理文本时高度关注分隔符。
(注:原文提供了图12,13,14,此处保留其引用)
图12 Llama-3-8B-Instruct中的注意力图示例(第0层和第0头)。
图13 Llama-3-8B-Instruct中的注意力图示例(第1层和第0头)。
图14 Llama-3-8B-Instruct中的注意力图示例(第2层和第0头)。

附录B:KV Cache的演变

图7详细说明了流式设置中KV Cache的演变。可以看出,在 $m_0$ 个token之后,$\pmb{n}$ 和 $Size_{run}$ 都是周期函数。且平均KV Cache使用量远小于最大容量 $\pmb{c}$。
图7 流式设置中KV Cache的演变。

附录D:不同模型的泛化性能

  • 不同架构与规模:将SepLLM适配到Pythia-6.9B、Pythia-12B和Falcon-40B。结果表明,对于规模相似的模型,设置相似的KV保留率可以产生同样良好的性能。较大的模型会有较低的困惑度,但需要更长的推理时间。
  • Base或Instruct模型:在LongAlpaca数据集上微调Llama-3-8B-Instruct和Base模型。发现两者均表现出优异性能(达到或超过原始注意力机制的模型)。区别在于Base模型需要微调更多步数。这表明Instruct模型的嵌入更好地对齐了SepLLM架构所需的分布。

附录F:大海捞针(Needle in a Haystack)测试

为了评估长上下文信息检索能力,进行了大海捞针实验。图8至11显示,即使丢弃了针(needle)中token的KV(除了可能存在的分隔符),SepLLM依然能在大多数场景下成功检索出针,而StreamingLLM无法完成此任务。这验证了SepLLM能有效地将片段信息压缩到分隔符的KV中。
图8 基于Pythia-160M-deduped的StreamingLLM的大海捞针测试结果。
图9 基于Pythia-160M-deduped的SepLLM的大海捞针测试结果。
图10 基于Llama-3-8B-instruct的SepLLM的大海捞针测试结果(保留4个初始token)。
图11 基于Llama-3-8B-instruct的SepLLM的大海捞针测试结果(保留32个初始token)。
图12 大海捞针图例说明

附录G与H:关于分隔符的选择与讨论

  • 分隔符的选择:将分隔符的选择视为一种超参数。实验证明,包含所有9种常见分隔符时性能最好。如果仅使用“.”和“?”,模型的推理能力会显著下降。
  • 为什么分隔符能保持性能
    • 从头训练视角:强制当前token只能看到相邻、分隔符和初始token,迫使模型通过自注意力机制将每个片段的信息浓缩到分隔符的KV对中。因此,分隔符的隐藏嵌入在功能上类似于RNN的状态空间。
    • 免训练视角:逗号、句号等是极高频的token。在预训练过程中,它们是词表中所有其他token最常见的上下文,因此它们的嵌入具有更高的相似性,在与其他token相乘时会产生更大的注意力值。从语义角度看,生成一个分隔符起到了总结当前片段的作用。

参考文献引用说明

  • 【13,Efficient Streaming Language Models with Attention Sinks+2024+ICLR+https://arxiv.org/abs/2309.17453】:在描述初始Token的作用 (Attention Sinks)、相邻Token的必要性、流式场景的挑战以及位置编码偏移策略时被引用,证明了保留初始和局部token对长序列生成的重要性。
  • 【25,H2O: HeavyHitter Oracle for Efficient Generative Inference of Large Language Models+2023+NeurIPS】:在描述相邻Token的必要性时被引用,作为保留局部窗口注意力以捕获局部依赖关系的支撑性文献。