Efficient Streaming Language Models with Attention Sinks

model 2309.17453
attention-sinkstreamingkv-cachelength-extrapolationsoftmax

StreamingLLM:用注意力汇(Attention Sink)实现无限长流式推理 #

1. TL;DR #

自回归 LLM 会把大量注意力"倾倒"到最初几个 token(注意力汇),因为 SoftMax 强制分母求和为 1。窗口注意力一旦驱逐这几个初始 token 就崩溃。StreamingLLM 保留 4 个初始 token 的 KV + 滑动窗口 KV,无需微调即可稳定处理 4M+ token,比重算基线快 22.2×。


2. Q1 / Q2 / Q3 #

Q1 — 痛点(Problem) #

把 LLM 部署到流式场景(多轮对话、日级长交互)有两个硬约束:

  1. KV 显存与延迟随序列线性增长 — 密集注意力缓存所有历史 token 的 K/V,长交互下显存爆掉、解码延迟单调上升(Figure 1a)。
  2. 长度外推能力有限 — 模型在超过预训练窗口(Llama-2 为 4K)后困惑度急剧退化。
  3. 最直观的补救——窗口注意力(只缓存最近 $L$ 个 token)——理论上显存/速度恒定,但论文证明它"一旦序列长度超过缓存大小、即使只驱逐第一个 token 的 KV,模型就立即崩溃"(Table 1:Llama-2-13B 的 PPL 从 5.40 暴涨到 5158.07)。唯一质量可接受的基线是滑动窗口 + 重算,但其窗口内二次注意力使复杂度为 $O(TL^2)$,慢到无法实用。

    Q2 — 方法(Method) #

    核心洞察:窗口注意力崩溃不是因为丢了近处 token,而是丢了最初几个 token。这些初始 token 攫取了异常高的注意力分数(与语义无关),论文命名为 attention sink(注意力汇)。成因是 SoftMax 归一化(Equation 1):即便当前 query 与任何历史 token 都不强匹配,注意力权重仍必须求和为 1,模型必须把这些"多余"的注意力倾倒到某处——初始 token 因对所有后续 token 可见而被训练成天然的倾倒点。

    StreamingLLM 的做法极简:KV cache = [4 个固定的初始 sink token 的 KV] + [滚动的最近窗口 KV],且位置编码按缓存内位置而非原文位置分配。无需任何微调即可恢复到接近"重算 oracle"的困惑度。第二贡献:从头预训练时预置一个可学习 sink token,可让单个 token 独揽注意力汇职责,之后流式部署只需这一个 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])。对 RoPE 需在旋转变换之前缓存 Key、每步解码时再施加位置变换;对 ALiBi 则施加连续线性偏置而非"跳跃式"偏置。忽略这一点,方法完全失效。

    Q3 — 结果(Results) #

    • 稳定性:Llama-2/MPT/Falcon/Pythia 全家族、[7,13,70]B 等各规模,在 4M+ token 上困惑度保持平稳(Figure 5)。
    • 匹配 oracle:在 20K token 上困惑度几乎与"滑动窗口+重算"基线持平(Figure 3)。
    • 效率:单 A6000 上,逐 token 解码延迟对缓存大小呈线性(基线呈二次),最高 22.2× 加速,显存与重算基线相当(Figure 10)。
    • 下游:流式 ARC-QA 与 StreamEval(120K token)准确率接近单样本 one-shot 基线;密集注意力 OOM、窗口注意力准确率近 0(Table 5)。

    3. 架构 / 方法图 #

    StreamingLLM 与四种方案的对比是全文的骨架图:

    Figure 1: StreamingLLM vs dense / window / recompute

    Paper's Figure 1, verbatim(caption: "Illustration of StreamingLLM vs. existing methods... (a) Dense Attention has $O(T^2)$... (b) Window Attention caches the most recent $L$ tokens' KV... performance declines sharply once the starting tokens' keys and values are evicted. (c) Sliding Window with Re-computation... $O(TL^2)$ complexity... (d) StreamingLLM keeps the attention sink (several initial tokens) for stable attention computation, combined with the recent tokens.")

    四象限清楚展示了 trade-off:(a) 质量随长度退化且显存无限增长;(b) 高效但一旦驱逐初始 token 就崩;(c) 质量好但二次复杂度太慢;(d) StreamingLLM 兼得效率与稳定。注意 (b) 与 (d) 的唯一区别就是那几个被保留的初始 token——这正是全文论点的可视化。

    StreamingLLM 的 KV cache 结构(两段式):

    Figure 4: StreamingLLM KV cache 结构

    Paper's Figure 4, verbatim(caption: "The KV cache of StreamingLLM.")

    缓存被划分为两部分:(1) Attention sinks(4 个初始 token)用于稳定注意力分布;(2) Rolling KV Cache 保留对语言建模至关重要的最近 token。读者应注意:sink 段是冻结的(永不驱逐),rolling 段随解码前移,位置编码在两段拼接后按缓存内下标重排。

    下面的 Mermaid 补充了原图未显式画出的逐步解码时的位置重分配逻辑(原图为静态快照):

    flowchart LR subgraph Cache["KV Cache (缓存内位置)"] S0["sink[0] pos=0"] --> S1["sink[1] pos=1"] --> S2["sink[2] pos=2"] --> S3["sink[3] pos=3"] S3 --> R0["recent pos=4"] --> R1["recent pos=5"] --> R2["recent pos=6"] --> R3["recent pos=7"] end Q["新 query (decode 第9个 token)"] -.attend.-> S0 Q -.attend.-> R3 R3 -->|驱逐最旧 recent, 前移| R2

    要点:原文序列位置可能是 [0,1,2,3,6,7,8](中间已被驱逐),但赋给注意力的是连续的 [0,1,2,3,4,5,6,7]——距离按缓存内而非原文计算。


    4. 作者证明 #

    本文无形式化定理,采用"机制论证 + 从头预训练验证"。两个 load-bearing 方程构成机制解释。

    记号表 #

    符号含义
    $x_i$token $i$ 的注意力 logit(pre-softmax 分数)
    $x_1$初始/第一个 token 的 logit
    $N$上下文 token 数
    $L$预训练注意力窗口大小
    $T$当前生成位置($T \gg L$)

    方程物理意义 #

    Equation 1(SoftMax 分母被首 token 主导)

    $$\text{SoftMax}(x)_{i}=\frac{e^{x_{i}}}{e^{x_{1}}+\sum_{j=2}^{N}e^{x_{j}}},\quad x_{1}\gg x_{j}$$

    把 $e^{x_1}$ 从求和中拎出,强调它对分母的主导。由于 $x_1 \gg x_j$,首 token 贡献了一大块近似恒定的归一化质量。驱逐它 = 抽走分母的大部分,导致所有其余 token 的注意力权重被整体放大、分布剧变——这就是窗口注意力崩溃的机制根因。

    Equation 2(SoftMax-off-by-one / Zero Sink)

    $$\text{SoftMax}_{1}(x)_{i}=\frac{e^{x_{i}}}{1+\sum_{j=1}^{N}e^{x_{j}}}$$

    分母额外的 $+1$ 等价于隐式预置一个 logit 为 0($e^0=1$)的 token,即"全零 K/V 的 token"。它给模型一个恒定的注意力倾倒口,使注意力总和不必强制为 1,从而不必劫持真实 token 当 sink。

    6 项最小核查 #

    1. 量纲/守恒:Eq.1 分母是全部 $e^{x_j}$ 之和,SoftMax 输出求和为 1,守恒成立。
    2. 极限行为:当 $x_1 \to \infty$,$\text{SoftMax}(x)_1 \to 1$,其余 $\to 0$——首 token 独占注意力,与 Figure 12"首 token 常获 >50% 注意力"定量吻合。
    3. 退化情形:若所有 $x_j$ 相等(无 sink),Eq.1 退化为均匀分布,去掉任一 token 影响均等——与"vanilla 模型需要多个初始 token"一致(Table 2:1 或 2 个不够、4 个才够)。
    4. Zero Sink 一致性:Eq.2 的 $+1$ 恒定项在 $x_j$ 都很小时占主导,把多余注意力吸走;但它不可学习,故 Table 3 显示 Zero Sink(0+1024 PPL=29214)远逊于 Learnable Sink(1235)。
    5. 绝对位置 vs 语义:Table 1 用 "\n" 替换前 4 个 token 仍恢复 PPL(4"\n"+1020 = 5.60 ≈ 4+1020 = 5.40),证明 sink 由绝对位置而非语义决定,符合"初始 token 对所有后续可见故易被训成 sink"的论证。
    6. 容量预算(KV bytes/token):StreamingLLM 缓存恒为 4 + L 个 token,故显存 $O(L)$ 恒定,与 Figure 10 观测到的"延迟线性、显存与重算相当"一致;密集注意力则 $O(T)$ 增长直至 OOM(Table 5)。

    7. 5. 实验与数据 #

      注意力汇现象的直接可视化(动机) #

      Figure 2: Llama-2-7B 注意力 logits 可视化

      Paper's Figure 2, verbatim(caption: "Visualization of the average attention logits in Llama-2-7B over 256 sentences... (1) The attention maps in the first two layers exhibit the 'local' pattern... (2) Beyond the bottom two layers, the model heavily attends to the initial token across all layers and heads.")

      这是全文的实证起点:除最底两层是"局部"模式外,几乎所有层/头都把大量注意力压在初始 token 上,无关其语义。这直接支撑 Q2 的机制论证——注意力汇是普遍存在的结构性现象,而非个例。

      稳定性对比:三种基线 vs StreamingLLM #

      Figure 3: 20K token 上各方法困惑度

      Paper's Figure 3, verbatim(caption: "Language modeling perplexity on texts with 20K tokens across various LLM... (1) Dense attention fails once the input length surpasses the pre-training attention window size. (2) Window attention collapses once... initial tokens are evicted. (3) StreamingLLM demonstrates stable performance, with its perplexity nearly matching that of the sliding window with re-computation baseline.")

      三条曲线的分岔点极具说服力:密集注意力在超过预训练窗口处翻车,窗口注意力在超过缓存大小处翻车,唯有 StreamingLLM 全程平稳并贴合 oracle(重算)。

      4M token 极长文本稳定性 #

      Figure 5: 4M token 困惑度(全家族全规模)

      Paper's Figure 5, verbatim(caption: "Language modeling perplexity of StreamingLLM on super long texts with 4 million tokens across various LLM families and scales. The perplexity remains stable throughout... with perplexity fluctuations due to book transitions.")

      跨 Llama-2/Falcon/Pythia/MPT 全家族、跨规模,困惑度在 4M token 上无漂移(波动仅来自 PG19 换书边界)。这是"无限流式"主张的核心证据。

      效率:22.2× 加速 #

      Figure 10: 解码延迟与显存对比

      Paper's Figure 10, verbatim(caption: "Comparison of per-token decoding latency and memory usage between the sliding window approach with re-computation baseline and StreamingLLM... StreamingLLM delivers a remarkable speedup of up to 22.2$\times$ per token and retains a memory footprint similar to the re-computation baseline.")

      关键在斜率:StreamingLLM 延迟随缓存大小线性增长,重算基线呈二次增长,二者之差随缓存增大而放大,故"最高 22.2×"是在大缓存端取得。显存两者相当,说明加速不以显存为代价。

      关键数值表 #

      Table 1(窗口注意力崩溃 → 4 个 sink 恢复,Llama-2-13B):

      ConfigPPL (↓)
      0+1024(窗口)5158.07
      4+10205.40
      4"\n"+10205.60

      Table 3(从头预训练 160M:vanilla / Zero Sink / Learnable Sink):

      Cache Config0+10241+10232+10224+1020
      Vanilla27.8718.4918.0518.05
      Zero Sink2921419.9018.2718.01
      Learnable Sink123518.0118.0118.02

      Learnable Sink 仅用 1 个 token(1+1023)就达 18.01,而 vanilla 需 4 个初始 token 才稳定——这是第二贡献的核心证据。

      Table 6(缓存大小非单调,反直觉):

      Cache4+2524+5084+10204+2044
      Falcon-7B13.6112.8412.3412.84
      MPT-7B14.1214.2514.3314.99

      增大缓存并不总降低困惑度(Falcon 4+2044 反而比 4+1020 差),暴露模型未能充分利用全部上下文的局限。


      6. 论证链 #

      步骤论证依据
      1流式部署需恒定显存/延迟,密集注意力做不到§1 Figure 1a;Table 5 Dense OOM
      2窗口注意力理论上恒定,但一旦驱逐初始 token 即崩溃Table 1(PPL 5.40→5158.07)
      3崩溃根因是初始 token 攫取了大量注意力(attention sink),驱逐它抽走 SoftMax 分母主体Figure 2;Equation 1;Figure 12(首 token >50% 注意力)
      4sink 由绝对位置而非语义决定("\n" 替换仍恢复)Table 1(4"\n"+1020=5.60)
      5因此只需保留 4 个初始 token + 滚动窗口,并按缓存内位置编码,即可稳定流式,无需微调Figure 3/4;Figure 5(4M token)
      6从头预训练预置 1 个可学习 sink token,可将 sink 收敛到单一 token,不损下游、略降困惑度Table 3/4;Figure 6/7
      7由此在流式 QA 与效率上均达标(接近 one-shot、22.2× 加速)Table 5;Figure 9/10

      7. 实现 cross-reference #

      官方代码:https://github.com/mit-han-lab/streaming-llm(本节维度/行号引用需以该 repo 实际版本核对;未逐行核实处标注)。

      • 两段式 KV cache(sink + rolling):核心逻辑在 streaming_llm/kv_cache.pyStartRecentKVCache(保留 start_size 个起始 token + recent_size 个近期 token)。构造缓存时按 [0:start_size][-recent_size:] 切片拼接。[实现细节以 repo 为准]
      • RoPE 缓存前置:论文 §3.2 指出"cache the Keys of tokens prior to introducing the rotary transformation",对应 repo 中对 Llama 的 pos_shift/enable_streaming_llm 补丁——在缓存时存旋转前的 Key、解码时按缓存内位置重新施加 RoPE(streaming_llm/pos_shift/modify_llama.py)。[实现细节以 repo 为准]
      • ALiBi 连续偏置:MPT 路径施加连续线性偏置而非跳跃偏置(对应 MPT 的 pos_shift 补丁)。[实现细节以 repo 为准]

      核心技术壁垒(§7 专段) #

      最难复现的一点是 "位置按缓存内下标而非原文下标分配""RoPE 在缓存后、解码时才施加旋转" 的组合。若照搬普通 KV cache(存旋转后的 Key、用原文绝对位置),当近期窗口滑动、原文位置出现跳变(如 [0,1,2,3,6,7,8])时,相对距离会"跳跃",注意力分布错乱,方法失效。论文用 [0,1,2,3,4,5,6,7] 的连续重编号消除这一跳变——这是把"保留 sink"从概念变为可用系统的关键工程细节(§3.2 worked example)。

      关键实现细节(1-2 个易漏点) #

      1. sink token 数 = 4 是经验阈值:1 或 2 个不足以恢复(Table 2),4 个够用、更多边际收益递减。这是默认配置,漏设会静默退化。
      2. Zero Sink ≠ Learnable Sink:SoftMax-off-by-one(隐式全零 token)只能部分缓解(Table 3 的 0+1024 仍是 29214 PPL),要彻底解决须在预训练预置可学习 sink token;把二者混为一谈会得到错误结论。

      3. 8. 部署与开放问题(LLM-specific) #

        服务化考量 #

        • 显存/并发:缓存恒为 4 + L 个 token 的 KV,与序列总长 $T$ 无关,故并发上限由 $L$(而非交互时长)决定——这正是"无限流式"在工程上的意义。
        • prefill vs decode:StreamingLLM 面向长交互的 decode 阶段(逐 token),其线性延迟优势在大缓存端才显著(Figure 10);短交互下相对重算优势有限。
        • 连续批处理友好性:缓存形状固定(sink 段冻结 + rolling 段定长),对静态 shape 的推理引擎友好,已被 TensorRT-LLM、HF Transformers、MLC LLM、Intel Extension for Transformers 采纳(Impact Statement)。

        明确的能力边界(不要误解"无限") #

        "无限流式" "无限上下文/记忆"。Table 7(§C)显示:一旦 query 与 answer 的 token 距离超过缓存大小,准确率降到 0;Table 8(§D)显示在依赖开头 prompt 的 LongBench 任务上,4+3496 配置逊于截断基线(需把 sink 数对齐到 1750 才追平)。因此 StreamingLLM 适合日常对话/短文档 QA,不适合长文档 QA/摘要。

        开放问题 #

        • 优势何时饱和/反转:Table 6 已显示缓存增大不单调降 PPL;模型对"缓存内长上下文"的利用率是瓶颈。
        • 跨模态可迁移性:现象已在 encoder BERT([SEP] 当 sink,Figure 14)与 ViT("registers")观察到,暗示 SoftMax 归一化而非自回归是根因;但 §I 显示 LLM 多加 sink token 无益甚至有害(Table 10:+2 sink 在 1+1023 = 25.73),与 ViT"多 register 有益"相反——跨模态的最优 sink 数并不一致。

        Appendix: 方法架构图(机制驱动) #

        代码来源:https://github.com/mit-han-lab/streaming-llm

        本文非常规"模型发布"(是训练无关的推理机制 + 一个 160M 验证性预训练),故按机制而非新架构组织。维度以 config/代码为准,未逐行核实处标注。

        A1 — Top-Level 数据流(不改架构,只改 KV cache 管理)

        flowchart TB T["token stream (T ≫ L)"] --> E["embed"] E --> B["N × 标准 Decoder Block (未改动)"] B --> KV{"KV cache 管理 (StreamingLLM 改动点)"} KV --> N["final norm"] --> H["lm_head → logits"] KV -.只此处替换.-> KVdetail["见 A2"]

        A2 — KV cache(sink + rolling)内部数据流

        flowchart LR new["新 token K/V"] --> append["append 到 rolling 段"] subgraph Cache sink["Sink 段: start_size=4 (冻结)"] roll["Rolling 段: recent_size=L"] end append --> roll roll -->|超出 recent_size| evict["驱逐最旧 recent (sink 永不驱逐)"] Cache --> repos["按缓存内下标重排位置 [0..len-1]"] repos --> rope["RoPE: 解码时施加旋转 / ALiBi: 连续线性偏置"] rope --> attn["注意力计算"]

        A3 — 主注意力变体(标准 MHA/GQA,唯一改动是 KV 来源与位置):Q/K/V 投影维度沿用宿主模型(Llama-2/Falcon/Pythia 用 RoPE,MPT 用 ALiBi),StreamingLLM 不改投影,只改"K/V 取自 [sink+rolling] 缓存"与"位置按缓存内下标"。

        A4 — 辅注意力变体N/A — 模型仅使用宿主模型原生注意力,无第二种 attention

        A5 — 选择/索引机制N/A — 无 Indexer/Routing;sink 段选择是固定前 4 个 token 的静态规则,非可学习门控(预训练版则为固定预置 1 个可学习 sink token)。

        A6 — 残差/连接机制N/A — 沿用标准 pre-norm 残差,未改动

        代码-图对照表 #

        代码构件对应图关键实现细节
        StartRecentKVCachekv_cache.pyA2start_size(=4) + recent_size(=L) 两段切片拼接
        modify_llama.py pos_shiftA2/A3旋转前 Key,解码时按缓存内位置施加 RoPE
        MPT pos_shift(ALiBi)A2施加连续线性偏置而非跳跃偏置
        预训练 sink tokenA2所有训练样本前置 1 个可学习 token(§3.3,Table 3)

        维度差异:论文用 4 个初始 token 作为默认 sink(vanilla 模型);预训练版仅需 1 个可学习 sink token(Table 3)。若代码默认 start_size 与此不符,以运行配置为准。