滑动窗口注意力(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 起就长期面对三件事:
- 表达力受限:需要把 exp(qkᵀ) 拆成 φ(q)φ(k)ᵀ,再加 kernel(Hedgehog、cosFormer 等)。
- 遗忘/覆写:固定大小的 state 必须持续重写,难保留长程关键信息。
- 训练昂贵:从零训练线性注意力 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 × wtoken 的远距依赖仍可能丢失。对真正超长上下文(≥ 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 注意力,没有根本解决二次方训练成本。
对工程落地的启发
- 第一动作:任何正在评估"是否要把 LLM 改造成 Linear Attention" 的团队,应先把 SWA(w, s) 当作训练免费的 baseline。SWA 在 0 训练成本下已经打平 Llamba 8-12B tokens 训练的工作,意味着 linearize 的 ROI 需要重新核算。
- 部署侧落地:SWA 推理时改 mask 不改权重,可以做成 vLLM / TensorRT-LLM 的 attention plugin。配合 attention sink(首 4 token 永远保留),不需要任何专用线性 kernel——FlashAttention / PagedAttention 都能直接受益。
- 组合策略:对需要更长上下文的场景,可考虑 Hybrid Attention(局部 SWA + 偶发 Full Attention 全局 token),或 SWA + 显式 memory tokens——这是本文 SWA 与 Linear Attention 都没覆盖的下一块工程机会。
- 不要盲目追新 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)
事实核查
- ✅ Table 1 数据可信度较高:表格数据与 alphaXiv/alphaxiv.org 摘要描述一致(MMLU 93.2 / 6-task avg 99.0),且覆盖 11 种 baseline 方法,数据源直接可追溯。但建议 v2 补一次对 alphaXiv 原文 Table 1 的 fetch 截图做锚点。
- ⚠️ "Microsoft(ASG)"归属需 fetch 原文确认:alphaXiv 摘要描述为"Microsoft"但未明确"ASG"(Applied Sciences Group)这一具体团队。"ASG"属于内部组织代号,可能来自作者 affiliation 而非论文正文自称——应在原文加 ⚠️ 说明"Microsoft ASG 来源待 PDF 首页核验",不宜作为确定性事实写入正文。
- ⚠️ "2-10× 量级差距"未给具体数字:原文写"BABILong / NIAH 上 SWA 比 Linear 高 2-10×",但具体是 NIAH 哪个长度(4K/32K/128K?)和哪个 Linear Attention 方法(Mamba / GLA / Llamba?)的对比均未明确,解读不应放大此数字的精确性——应改为"⚠️ NIAH / BABILong 具体倍数待 PDF §5 核验,2-10× 为原文定性描述,非精确数字"。
- ✅ SWA 公式与伪代码一致:公式描述与伪代码实现逻辑吻合(窗口 w + s 个 sink),没有内部冲突。
- ⚠️ QRWKV7 / QLinAtt / ARWKV 等论文未核实:这些是 2025 年新工作,具体出处(会议/期刊)未核实,不应在解读中传播这些工作"已发表"的误读——建议降级为"Zhang'25a / Goldstein'25 / Yueyu'25"等未核实的笼统标注,或直接删去。
工程落地与实践
-
vLLM 接入 SWA 的两种路径: - 路径 A(推荐,快): 用
vllm/attention.py的Attention.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 确定瓶颈在哪。 -
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 没有被意外展开。 -
实际吞吐量 / 延迟实测预估: - 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 的恢复率数字推断工程收益。 -
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。 -
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。
-
vs FlashAttention-3 / cuDNN flash attention: - SWA 的 O(w) attention 计算仍然可以用 FlashAttention 加速(FlashAttention 支持 arbitrary mask),不需要替换 kernel; - 真正值得评估的是:FlashAttention 的"streaming attention mode"(每步只算局部 w)是否与 SWA 等价——如果等价,则 SWA 的额外工程工作价值有限。