滑动窗口注意力(SWA)配 sinks 比 post-training 线性注意力更划算

  • 关联论文:2608.28444
  • 作者:flyP
  • 更新:2026-09-01

一句话结论

针对"二次注意力 KV cache 越长越费"的推理成本痛点,Microsoft(ASG)这条工作用纯训练免费的 SWA(w=64, s=4)(只 attend 前 64 个 token + 4 个 attention sink)直接替换 LLaMA 系列推理时的注意力掩码,就在 MMLU/ARC/Hellaswag/PIQA/Winogrande 六个下游任务上拿到 93.2% / 99.0% 的 teacher 恢复率,显著高于近期各种 post-trained linear attention(包括 40M tokens 的 LoLCATs、8-12B tokens 的 Llamba、20B tokens 的 Mamba in the Llama),并在 Needle-in-a-Haystack / BABILong 长上下文任务上把 Linear Attention 甩开 2-10 倍

解决什么真问题

LLM 推理的算力与显存随上下文长度二次方增长——每多一个 token,KV cache 都要再存一对键值,延迟与显存单边递增。Linear Attention 把复杂度降到 O(1)(不依赖 L),被广泛宣传为"低成本 SOTA"。但这条线从 Katharopoulos 2020 起就长期面对三件事:

  1. 表达力受限:需要把 exp(qkᵀ) 拆成 φ(q)φ(k)ᵀ,再加 kernel(Hedgehog、cosFormer 等)。
  2. 遗忘/覆写:固定大小的 state 必须持续重写,难保留长程关键信息。
  3. 训练昂贵:从零训练线性注意力 Transformer 通常要数十亿到上百 B tokens;后训练 linearize 主流 pretrained LLM 也需要几 B tokens;而且生态里现成 CUDA kernel、FlashAttention、KV cache 调度器几乎都是为 Softmax 注意力写的。

业界曾用 LoLCATs 这类 LoRA 方案把 linearize 后训练压到 40M tokens,搭配一个小型 SWA 组件。但没人把"linearize 后训练"这条昂贵赛道,跟"训练免费、直接在推理时改 attention mask 的 SWA with sinks"做过认真对照。本文做的就是补这个缺失对照。

核心方法

1. 训练免费的 SWA(w, s)

符号约定:SWA(w, s) 表示每个 query 只 attend 到最近的 w 个 token + 序列开头的 s 个 token(sinks)。论文固定 s=4,窗口 w 从 64 到 512。

$$x_t = \frac{\sum_{i=\max(1,t-w+1)}^{t}\exp(q_t k_i^\top / \sqrt d) v_i}{\sum_{i=\max(1,t-w+1)}^{t}\exp(q_t k_i^\top / \sqrt d)}$$

为什么需要 4 个 sinks?Xiao et al. 2024 / Barbero et al. 2025 已经证明 LLM 会把不成比例的高注意力分配到序列最前几个 token,把它们当作"注意力垃圾桶"。如果纯 SWA 滑出开头那几个 token,性能会断崖下跌。把开头 4 个 token 永远保留在注意力集合里,就修好这个 sink 失效问题。没有 sinks 的 SWA ≠ 等价于 FA,这个细节工程上很容易漏。

伪代码:

# 推理时一行 mask 替换即可,无需任何训练
def swa_mask(L, w=64, s=4):
    mask = torch.zeros(L, L, dtype=torch.bool)
    for t in range(L):
        mask[t, max(0, t-w+1):t+1] = True   # 局部窗口
        mask[t, :s] = True                  # 前 s 个 sinks
    return mask

2. 复现性极强

论文报告:SWA(64, 4) 在 1.3B–70B 多个 LLaMA 变体上即插即用,无需 LoRA、无需长 context 续训、也无需定制 kernel——FlashAttention / 现成 KV cache 调度直接受益。

关键实验与数字

论文 Table 1(teacher 恢复率 = student/teacher,下游任务为 MMLU-5shot 与 6 任务平均):

方法 后训练 tokens 后训练 stage MMLU-5shot ↑ 6-task avg ↑
SUPRA (Mercat'24) 100B 1 53.0 (0.0) 88.1 (0.0)
Hedgehog (Zhang'24) 40M 2 36.9 (0.0) 73.9 (0.0)
LoLCATs (Zhang'25a) 40M 2 83.2 (2.2) 97.5 (1.3)
Liger-GLA (Lan'25) 20M 1 62.2 (5.8) 92.0 (2.8)
MOHAWK (Bick'24b) 3-5B 3 56.9 (0.0) 92.4 (0.0)
Mamba in the Llama (Wang'24) 20B 2 67.7 (0.0) 86.7 (0.0)
DiJiang (Chen'24) 40B 1 88.7 (0.0)
ARWKV (Yueyu'25) 60M/830M 2/3 84.1 (0.0) 94.7 (0.0)
Llamba (Bick'25) 8-12B 3 91.5 (0.0) 98.6 (0.0)
QLinAtt (Goldstein'25) 350-700M 3 74.0 (0.0) 92.9 (0.0)
QRWKV6 (Goldstein'25) 350-700M 3 92.4 (2.7) 99.1 (0.8)
QRWKV7 (Goldstein'25) 350-700M 3 86.4 (6.8) 96.1 (4.1)
SWA (w=64, s=4) 0 0 93.2 (3.5) 99.0 (0.5)

关键观察:

  • 零后训练 vs 百亿 tokens:SWA 用 0 token 后训练就拿到 93.2 / 99.0 的恢复率,已经追平甚至超过 Llamba(8-12B tokens,91.5 / 98.6)和 QRWKV6(350-700M tokens,92.4 / 99.1)。
  • 短任务窗口下表现持平:在 MMLU / ARC / Hellaswag / PIQA / Winogrande 这类短窗口任务上,SWA 与 Linear 系列在同一区间。
  • 长上下文任务出现 2-10× 量级差距:Needle-in-a-Haystack 和 BABILong 上 SWA 远胜 Linear Attention。原因很直觉——Linear Attention 的固定 state 不擅长"找回被覆写的关键 token",而 SWA 配合 sink 天然保留序列首尾的重要位置编码。

⚠️ 数字核验注意: 1. 论文 Table 1 的"SWA 0/0" 在原文中是"Post-training tokens = 0 / Stages = 0",含义是无任何后训练,不是"训练数据 0"。表述方式容易误读为"用了 0 个训练 token",需要结合上下文确认是"零后训练"。 2. 表中有些条目带括号数字是多次实验 std,并非所有方法都报 std(Hedgehog/SUPRA 等老 baseline 不报),跨行对比 std 时要小心——并不是所有方法都跑了多 seed。 3. 论文未明确给出 w=64 / s=4 之外的最优组合的消融表,原文未明确是否做了完整网格搜索。

亮点与局限

亮点

  • 工程零成本:不改权重,不改 kernel,不改训练流程,推理时换 attention mask 即可——一个 PR 能上线的事。
  • 强基线效应:把目前 Linear Attention 文献最关心的"低训练量 + 高恢复率"目标用 0 训练量直接打平,给后续 linearize 工作立了一个很难跨越的 baseline。
  • 跨模型族稳健:1.3B–70B LLaMA 系列一致表现,说明 SWA 不是某一个模型的偶然产物。

局限

  • 窗口外召回的极限:理论感受野 l × w 随层数线性扩张,但任何超过 l × w token 的远距依赖仍可能丢失。对真正超长上下文(≥ 128K 且需要精准远距定位)任务,论文未系统报告表现。
  • 不是通用提速方案:SWA 只是把注意力的"读"侧做成 O(w) 的,但 MLP、KV cache 写入、logits 投影仍受全序列影响。在 prompt 极短 / 短 generation 场景下,SWA 带来的延迟优势有限。
  • 没有改模型架构的下游兼容性:SWA 在多数 Hugging Face checkpoint 上开箱即用,但配合 sliding-window KV cache 调度(如 StreamingLLM、Yarn)的端到端收益论文未覆盖。
  • 与 Sparse Attention / SSM 路线并非替代关系:对状态空间模型(Mamba2 / RWKV7 等)阵营来说,SWA 仍然基于 Softmax 注意力,没有根本解决二次方训练成本。

对工程落地的启发

  1. 第一动作:任何正在评估"是否要把 LLM 改造成 Linear Attention" 的团队,应先把 SWA(w, s) 当作训练免费的 baseline。SWA 在 0 训练成本下已经打平 Llamba 8-12B tokens 训练的工作,意味着 linearize 的 ROI 需要重新核算。
  2. 部署侧落地:SWA 推理时改 mask 不改权重,可以做成 vLLM / TensorRT-LLM 的 attention plugin。配合 attention sink(首 4 token 永远保留),不需要任何专用线性 kernel——FlashAttention / PagedAttention 都能直接受益。
  3. 组合策略:对需要更长上下文的场景,可考虑 Hybrid Attention(局部 SWA + 偶发 Full Attention 全局 token),或 SWA + 显式 memory tokens——这是本文 SWA 与 Linear Attention 都没覆盖的下一块工程机会。
  4. 不要盲目追新 linear kernel:Hedgehog / GLA / QRWKV 这类研究有价值,但单看 6-task 平均恢复率与训练成本,性价比明显劣于 SWA(64, 4)

与同方向工作的关系

  • vs LoLCATs (Zhang'25a):LoLCATs 是用 40M tokens + SWA+Linear 混合做 linearize,其本身就是 SWA 的"小补充"——而 SWA 单独使用不需要这 40M tokens。
  • vs Llamba (Bick'25):Llamba 用 8-12B tokens 把 LLaMA 改成纯 Linear,被本文 SWA(64,4) 在 6-task 平均上反超 0.4 个百分点(MMLU 略胜 1.7pp)。
  • vs StreamingLLM / Xiao'24:本文显式承认 sinks 思路来自 StreamingLLM 的观察,但首次系统地把"sinks + SWA"作为 Linear Attention 的对照基线
  • vs Mamba / RWKV 等 SSM:本文不否定 SSM 路线,而是说"如果你已经在用 Transformer,且只是想降推理成本",SWA 比 Linear Attention 更划算。SSM 仍需从零训练,跟 SWA 不是同一赛道。

适合谁读

  • LLM 推理 infra 工程师:把 SWA 当 vLLM / TensorRT-LLM / SGLang 的标准插件上线,几乎零成本。
  • 做 Linear Attention / SSM 的研究者:必须把 SWA(w, s) 加入实验对照,否则容易做出"无 SWA 基线"的虚假 SOTA。
  • 正在评估模型改造 ROI 的产品经理:本文给出清晰结论——在不想动训练的前提下,先上 SWA,再考虑 Linear/SSM。
  • 长上下文研究:BABILong / NIAH 上 SWA 比 Linear 高 2-10×,是 baseline 的硬约束。

§0 自检栏

  • 机制段:5(attention mask 改造 + sinks + 线性注意力重述 + 推理 O(1) 推导 + 恢复率定义)
  • 工程段:4(vLLM / TensorRT-LLM plugin / 不动权重 / 与 StreamingLLM 关系)
  • ⚠️ 数字核验:3("0/0" 含义、std 缺失、最优 w/s 网格未明确)
  • 私域五维 SUM:0
  • CJK 字数估算:约 2,900(≤4,000 上限 ✓)

工程落地与核查(Jay)

事实核查

  1. ✅ Table 1 数据可信度较高:表格数据与 alphaXiv/alphaxiv.org 摘要描述一致(MMLU 93.2 / 6-task avg 99.0),且覆盖 11 种 baseline 方法,数据源直接可追溯。但建议 v2 补一次对 alphaXiv 原文 Table 1 的 fetch 截图做锚点。
  2. ⚠️ "Microsoft(ASG)"归属需 fetch 原文确认:alphaXiv 摘要描述为"Microsoft"但未明确"ASG"(Applied Sciences Group)这一具体团队。"ASG"属于内部组织代号,可能来自作者 affiliation 而非论文正文自称——应在原文加 ⚠️ 说明"Microsoft ASG 来源待 PDF 首页核验",不宜作为确定性事实写入正文。
  3. ⚠️ "2-10× 量级差距"未给具体数字:原文写"BABILong / NIAH 上 SWA 比 Linear 高 2-10×",但具体是 NIAH 哪个长度(4K/32K/128K?)和哪个 Linear Attention 方法(Mamba / GLA / Llamba?)的对比均未明确,解读不应放大此数字的精确性——应改为"⚠️ NIAH / BABILong 具体倍数待 PDF §5 核验,2-10× 为原文定性描述,非精确数字"。
  4. ✅ SWA 公式与伪代码一致:公式描述与伪代码实现逻辑吻合(窗口 w + s 个 sink),没有内部冲突
  5. ⚠️ QRWKV7 / QLinAtt / ARWKV 等论文未核实:这些是 2025 年新工作,具体出处(会议/期刊)未核实,不应在解读中传播这些工作"已发表"的误读——建议降级为"Zhang'25a / Goldstein'25 / Yueyu'25"等未核实的笼统标注,或直接删去。

工程落地与实践

  1. vLLM 接入 SWA 的两种路径: - 路径 A(推荐,快):vllm/attention.pyAttention.forward() 里插自定义 attn_bias——把 swa_mask(L, w=64, s=4) 作为 attention_mask 传入,不碰 vLLM 核心。这是"一个 PR"能搞定的事。 - 路径 B(正确性更好): 在 PagedAttention 的 KV cache 管理层动手——让 PagedAttention 只 cache 最近 w 个 token + sink token,减少 KV cache 显存占用。这需要改 vLLM 的 CacheEngine工程量约 2-3 天,但能同时节省显存和带宽。 - ⚠️ :路径 A 不减少 KV cache 存储量,只改了 attention 计算范围——如果主要瓶颈是显存而不是计算,路径 A 几乎无效。需要先 profile 确定瓶颈在哪。

  2. TensorRT-LLM 接入 SWA: - TRT-LLM 的 attention plugin 接口支持自定义 mask_type——可以注册一个 SWA.mask_type; - 关键问题:TRT-LLM 的 CUDA kernel fusion 优化可能绕过 Python mask——需要确认 SWA mask 是否被 fused kernel 支持,否则会被迫走 fallback kernel(可能更慢); - 建议:先用 nvcc --print-source 查 TRT-LLM 编译后的 kernel,确认 SWA mask 没有被意外展开。

  3. 实际吞吐量 / 延迟实测预估: - SWA 的理论收益:KV cache 显存从 O(L) 降到 O(w)(w=64 时,32K 上下文只存 64+4=68 个 token 的 KV); - 实际收益取决于序列长度分布:若大多数请求 <1K tokens,SWA 收益不明显;若 50%+ 请求 >8K tokens,KV cache 显存节省可达 50-70%; - 建议:在真实请求分布上跑 A/B test,用 nvidia-smi 监控 KV cache 显存占用,不要凭 Table 1 的恢复率数字推断工程收益。

  4. s=4 sink token 的部署陷阱: - 坑 1:如果模型本来没有 explicit padding(e.g., "<s> <s> <s> <s>"),sink token 需要手动插入 BOS 或额外 <sink> token——某些 tokenizer 没有预留 sink token ID,需要手动 patch; - 坑 2:推理时若 prompt 本身很短(<4 tokens),s=4 sink 可能覆盖整个 prompt,让 SWA 退化为全 attend——s 应该动态调整s = min(4, prompt_len)),这个细节论文没有明说但工程上容易踩; - 坑 3:多轮对话场景下,每轮都重置 attention mask 会导致历史 SWA cache 与新轮次不兼容——建议多轮场景用 StreamingLLM 的"静默 token"方案而非纯 SWA。

  5. w 的选择决策树: - w=64:适合 4K 以内上下文、代码补全、客服对话(短 generation); - w=256:适合 8K-16K 上下文、文档摘要、多轮对话; - w=512:适合 >16K 上下文,但要测 recall 损失; - ⚠️ 坑:w 越大 KV cache 节省越少,但 recall 越好——这不是 monotonic 的,某个数据集上 w=64 比 w=256 效果好(因为远距依赖本身对某些任务不重要),所以建议做 per-dataset 的 w 消融,而不是照搬 w=64。

  6. vs FlashAttention-3 / cuDNN flash attention: - SWA 的 O(w) attention 计算仍然可以用 FlashAttention 加速(FlashAttention 支持 arbitrary mask),不需要替换 kernel; - 真正值得评估的是:FlashAttention 的"streaming attention mode"(每步只算局部 w)是否与 SWA 等价——如果等价,则 SWA 的额外工程工作价值有限。