AI 技术
#注意力沉没#StreamingLLM#KV缓存#流式推理#长文本生成

注意力沉没与流式大模型:如何让无限长文本生成保持稳定

流式大模型在超长对话中常因注意力崩溃而输出混乱,滑动窗口方案虽节省显存却会引发困惑度飙升。本文以对话机器人为场景,解释注意力沉没现象及其成因,分析 StreamingLLM 如何通过保留初始 token 的 KV 缓存来稳定注意力分布,并对比滑动窗口、长度外推等方法的权衡,帮助读者理解在无限长流式生成中保持模型稳定的工程决策。

当对话机器人开始胡言乱语

假设你正在与一个基于大模型的客服机器人进行多轮对话,前几十轮它都能准确理解你的问题并给出专业回答。但当对话进行到第 500 轮时,机器人的回复开始变得前言不搭后语,甚至完全脱离上下文。你检查了显存占用,发现它还在正常范围内,推理速度也没有明显下降。那么,是什么导致了模型的突然崩溃?

这个场景揭示了大模型在流式部署中的一个核心矛盾:自回归生成需要缓存所有历史 token 的 Key 和 Value 状态(KV 缓存),但显存是有限的。当对话长度超出训练时的注意力窗口大小时,模型不仅面临显存爆炸的风险,更会遭遇一种被称为“注意力沉没”的失效模式。直观的解决方案是只保留最近的一部分 token,即滑动窗口注意力,但实验表明,一旦窗口外的 token 被丢弃,模型的困惑度会急剧上升,即使这些 token 在语义上已经不再相关。

为什么滑动窗口会失败

滑动窗口注意力是一种自然的工程直觉:既然显存放不下所有 KV 缓存,那就只保留一个固定大小的窗口,覆盖最近的若干 token。这种方式确实能维持稳定的显存占用和解码速度,但在序列长度超过窗口大小后,模型性能会出现断崖式下跌。

图 3 展示了 Llama-2、MPT、Falcon 和 Pythia 等模型在 2 万 token 文本上的语言建模困惑度。当文本长度超过缓存大小时,由于初始 token 被排除,困惑度会激增。这表明,无论初始 token 与当前预测位置的距离有多远,它们对于保持 LLM 的稳定性都至关重要。

研究者通过可视化注意力分数揭示了背后的原因。图 2 展示了 Llama-2 7B 模型所有层和头的注意力分布:除了最下面两层之外,模型在所有层和头都始终对初始 token 分配了极高的注意力权重。这些初始 token 就像一个个“注意力沉没”(attention sink),吸收了大量的注意力分数,即使它们在语义上并不重要。

这一现象源于 Softmax 操作的数学约束。Softmax 要求所有注意力分数的总和为 1,因此当当前 query 与大多数先前 token 没有强匹配时,多余的注意力值必须分配到某些位置。由于自回归语言建模的特性,初始 token 对几乎所有后续 token 都是可见的,因此模型在训练过程中自然学会了将它们作为注意力沉没。一旦滑动窗口丢弃了这些初始 token,Softmax 分母中原本由它们占据的很大一部分就会消失,导致注意力分数的分布发生剧烈变化,偏离了模型在训练和正常推理中预期的分布。

注意力沉没的数学直觉

为了更精确地理解注意力沉没,我们可以从 Softmax 的计算公式入手。对于某个 query,其与第 i 个 key 的注意力分数为:

SoftMax(x)i=exij=1Nexj\text{SoftMax}(x)_i = \frac{e^{x_i}}{\sum_{j=1}^{N} e^{x_j}}

其中 xix_i 是 query 与第 i 个 key 的点积。当初始 token 的注意力分数 x1x_1 远大于其他 token 的分数 xjx_jj2,...,Nj \in 2,...,N)时,分母主要由 ex1e^{x_1} 贡献。如果移除这些初始 token,分母会大幅减小,导致剩余 token 的注意力权重被不成比例地放大,破坏了模型原本习得的注意力模式。

研究者通过实验进一步确认,初始 token 的重要性在于其绝对位置而非语义内容。当用换行符“\n”替换前四个 token 后,模型仍然会非常重视这些初始换行符。这说明模型学习到的是“将多余的注意力分配给序列开头的固定位置”,而不是依赖于特定的语义信息。

StreamingLLM:保留沉没与滑动窗口的结合

基于上述观察,研究者提出了 StreamingLLM,一种无需任何微调即可让有限注意力窗口的 LLM 泛化到无限序列长度的框架。其核心思想极其简洁:在 KV 缓存中同时保留注意力沉没(几个初始 token)和滑动窗口(最近的 token),而丢弃中间的大量 token。

图 4 展示了 StreamingLLM 的 KV 缓存结构,它在概念上分为两部分:

  • 注意力沉没:通常保留 4 个初始 token 的 KV 状态,用于稳定注意力计算中的 Softmax 分布。
  • 滑动 KV 缓存:保留最新的 token,这些 token 携带了最近的上下文信息,对语言建模至关重要。

在位置编码的处理上,StreamingLLM 有一个关键细节:它根据 token 在缓存中的相对位置来分配位置信息,而不是它们在原始文本中的绝对位置。例如,如果当前缓存包含 token [0, 1, 2, 3, 6, 7, 8],那么在解码第 9 个 token 时,分配的位置将是 [0, 1, 2, 3, 4, 5, 6, 7],而不是原始位置 [0, 1, 2, 3, 6, 7, 8, 9]。这种相对位置编码策略确保了模型始终在一个紧凑的、连续的位置空间中工作,避免了因位置跳跃而引入的分布外问题。

下面的流程图展示了 StreamingLLM 在流式对话中的核心决策过程:

flowchart TD
    A[新 token 生成] --> B{KV 缓存是否已满?}
    B -- 否 --> C[将新 token 的 KV 追加到缓存]
    B -- 是 --> D[丢弃中间 token 的 KV]
    D --> E[保留初始 4 个 token 的 KV]
    E --> F[保留最近 W 个 token 的 KV]
    F --> G[重新计算位置编码]
    G --> H[执行注意力计算]
    C --> H
    H --> I[输出下一个 token]
    I --> A

在对话机器人的场景中,当对话轮次不断增长,KV 缓存达到预设大小后,系统会丢弃窗口之外的历史 token,但始终保留对话开始时的前几个 token。这样,无论对话进行多长,模型都能维持稳定的注意力分布,避免困惑度飙升。

与滑动窗口、长度外推的对比

为了更清晰地定位 StreamingLLM 的工程价值,下表将其与滑动窗口、带重计算的滑动窗口以及长度外推方法进行了多维度对比。

方法显存占用推理速度长文本困惑度是否需要微调支持无限长度
密集注意力随序列长度线性增长逐渐变慢超过训练长度后失效
滑动窗口恒定超过窗口大小后激增
滑动窗口 + 重计算恒定慢(需重计算)稳定
长度外推(如 RoPE 插值)随序列长度增长逐渐变慢有限范围内稳定通常需要有限
StreamingLLM恒定稳定

滑动窗口的主要优势在于显存恒定和速度快,但无法处理超过窗口长度的上下文。带重计算的滑动窗口通过从最近的 token 重建 KV 状态来恢复性能,但计算开销巨大,在流式应用中不实用。长度外推方法(如位置插值)试图扩展模型的上下文窗口,但通常需要微调,且扩展范围有限,显存占用仍会随序列长度增长。

StreamingLLM 在显存、速度和困惑度之间取得了实用的平衡:显存占用恒定,推理速度与滑动窗口相当,且困惑度能与重计算基线相媲美。在流式设置下,StreamingLLM 相比滑动窗口重计算基线实现了高达 22.2 倍的加速。

预训练中的沉没 Token:从被动发现到主动设计

研究者进一步探索了在预训练阶段主动引入注意力沉没的可能性。既然模型在训练中会自发地将初始 token 作为沉没,那么能否通过设计一个专门的“沉没 Token”(Sink Token)来更优雅地解决这个问题?

实验表明,在预训练时于每个样本开头添加一个可学习的 Sink Token,可以让模型学会将所有多余的注意力集中到这一个 token 上。图 6 显示,使用 Sink Token 训练的模型与普通模型具有相似的收敛动态和性能,但普通模型需要多个初始 token 作为沉没才能保持稳定,而使用 Sink Token 的模型仅靠这一个 token 就能实现令人满意的效果。

图 7 的注意力可视化对比更直观地展示了差异:普通模型的注意力在深层会分散到多个初始 token 上,而使用 Sink Token 训练的模型在所有层和头都一致地聚焦于该 token,形成了更高效的注意力卸载机制。

另一种替代方案是修改 Softmax 函数本身,即“Zero Sink”方法:

SoftMax1(x)i=exi1+j=1Nexj\text{SoftMax}_1(x)_i = \frac{e^{x_i}}{1 + \sum_{j=1}^{N} e^{x_j}}

这个变体不要求所有上下文 token 的注意力分数总和为 1,相当于在注意力计算中引入了一个 KV 特征全为零的虚拟 token。这种方法与添加 Sink Token 在效果上等价,但无需额外参数。

工程部署中的权衡与观察

在真实对话机器人系统中部署 StreamingLLM 时,需要关注几个关键的工程维度。

显存与缓存管理:StreamingLLM 的 KV 缓存大小由窗口大小 W 和保留的初始 token 数 S 决定,总缓存大小为 S + W。在实际部署中,S 通常取 4,W 则根据显存预算和延迟要求设定。由于缓存大小恒定,系统可以预分配显存,避免动态分配带来的碎片和开销。

位置编码的一致性:如前所述,StreamingLLM 必须使用缓存中的相对位置进行编码。如果模型使用了 RoPE 等相对位置编码,这一过程是自然的;但如果模型依赖绝对位置编码,则需要确保位置重新映射的正确性,否则会引入分布偏移。

困惑度监控:在生产环境中,困惑度是检测注意力崩溃的关键指标。当困惑度出现异常尖峰时,通常意味着 KV 缓存中的注意力沉没被意外丢弃或位置编码计算错误。应当设置告警阈值,并在检测到异常时触发缓存重置或回退到更保守的窗口策略。

失败模式:StreamingLLM 并非万能。当对话的早期内容包含关键信息,而这些信息在缓存中已被丢弃时,模型将无法回忆起这些细节。例如,用户在对话开头提供了自己的姓名和订单号,经过数百轮对话后,这些信息可能已经不在滑动窗口内,模型将无法准确引用。这种“早期信息遗忘”是 StreamingLLM 的固有局限,需要在应用层通过外部记忆或摘要机制来补偿。

此外,如果初始 token 本身被污染(例如包含特殊字符或异常嵌入),注意力沉没可能会失效,因为模型无法将多余的注意力有效卸载到这些 token 上。在极端情况下,这可能导致整个注意力分布崩溃。

未解决的问题与适用边界

StreamingLLM 为流式大模型提供了一种实用的无限长度生成方案,但它并非解决所有长文本问题的银弹。

首先,该方法牺牲了模型对远距离信息的精确回忆能力。在需要跨长距离进行精确信息检索的任务中(如长文档问答),StreamingLLM 的表现会显著弱于密集注意力或带有检索增强的方案。

其次,注意力沉没的有效性依赖于模型在训练过程中确实学会了将初始 token 作为沉没。对于某些经过特殊微调或结构修改的模型,这一假设可能不成立,需要额外的适配工作。

最后,Sink Token 预训练虽然能提升 StreamingLLM 的性能,但需要修改训练流程,对于大多数基于现有模型进行部署的团队来说并不现实。Zero Sink 的 Softmax 变体提供了一种无需训练的替代方案,但其在更广泛模型和任务上的鲁棒性尚需更多验证。

在对话机器人的工程实践中,StreamingLLM 最适合那些对话轮次可能极长、但单轮信息价值相对较低的场景,如闲聊机器人、长期运行的监控 Agent 等。对于需要精确记忆早期对话细节的任务,应当结合外部记忆模块或定期摘要机制,将 StreamingLLM 作为底层流式推理引擎,而非完整的记忆解决方案。

资料来源

  1. Efficient Streaming Language Models with Attention Sinks
  2. StreamingLLM:一种能够接受近乎无限长度文本的大模型框架
  3. 關於 Transformer 架構中的Attention 機制