滑窗注意力 Sliding Window Attention

Sliding Window Attention / SWA

稀疏注意力中最常用、工程性价比最高的一种特例:规定每个 token 只能关注其自身以及向前追溯的最近 W 个相邻 token(W 为窗口大小,通常 4096 或 8192),实现 O(N·W)≈O(N) 的线性复杂度。

详细解释

Sliding Window Attention(SWA,滑窗注意力 / 滑动窗口注意力) 一句话定义:对于长度为 N 的上下文序列,第 i 个 token 的注意力不再计算 [0, N-1] 全部 N 个 key,只计算 [i-W+1, i] 区间内最近 W 个 key(含自身)。它是稀疏注意力家族里最简单、最鲁棒、几乎零训练成本就能直接上的那一个,Mistral 7B 首次把它用在 7B 级开源模型上后,2024 年之后所有新模型(Llama 3.1 128K、Qwen 2.5 128K、Gemma 2 等)全把它设为默认长上下文注意力机制了。

SWA 的三大核心优势让它碾压其他所有稀疏注意力变种:(1) 精度损失极小——在绝大多数自然语言/代码数据上,一个 token 真需要看的 95% 以上权重都集中在最近 ±2000 个相邻邻居里,切掉更远的几乎对下游任务没影响;(2) 工程实现极其简单——只需要在 Attention Mask 上做一个下三角带宽为 W 的带状掩码,不需要路由、不需要分块、不需要全局锚点,推理引擎 vLLM/SGLang 全原生支持,改一行配置 sliding_window=4096 就能开;(3) KV Cache 变成环形队列大幅省显存——全注意力的 KV Cache 长度随上下文线性增长,128K 上下文的 KV Cache 7B 模型就要吃 80G 显存;SWA 的 KV Cache 只要存最近 W 个 token,W=4096 时无论上下文是 128K 还是 1M,KV Cache 都只占 4096 的固定显存,长上下文推理成本直接砍 90%。

唯元智创 平台的所有开源长上下文模型默认都开了 SWA + FlashAttention-2 的组合:对于 32K 以上的请求自动切换到 SWA 模式,同时针对需要「看开头全局信息」的任务(如全文摘要、合同签约方抽取),我们额外加了一个「全局锚点重注入」的小 trick——把文档前 256 个 token 的 KV 永远保存在 Cache 最前面不被滑窗覆盖,既保留 SWA 的低成本又不丢开头信息。

滑窗注意力的两个核心工程参数(W 和步长 S)详解

参数定义典型值调参建议
窗口大小 W(Window Size)每个 token 能看的邻居数量(含自身)2048 / 4096 / 8192 / 16384中文/混合场景推荐 4096(甜点),代码场景推荐 8192(代码变量引用跨距更大),16384 以上收益递减极快不推荐
滑窗步长 / 重叠率 S(Stride / Overlap)多层叠加时不同层窗口的偏移量,S=1 不偏移、S=W/2 相邻层错开半个窗口S=1(Mistral 默认)或 S=W/2(Sliding Window with Dilated Rolling)不做超长文本理解 S=1 够用;要捕捉 W 以上的长程依赖可让不同层 S 依次偏移 0、W/4、W/2、3W/4,叠加后能以低成本覆盖更长范围

一个直观的 SWA 注意力可视化(W=4)

序列 = [t0, t1, t2, t3, t4, t5, t6],每个 ti 可看的 token(✅ 可看,❌ 被掩码忽略):

注意力矩阵行=query 列=key
    t0 t1 t2 t3 t4 t5 t6
t0  ✅ ❌ ❌ ❌ ❌ ❌ ❌  → 只能看自己(左边没有4个)
t1  ✅ ✅ ❌ ❌ ❌ ❌ ❌
t2  ✅ ✅ ✅ ❌ ❌ ❌ ❌
t3  ✅ ✅ ✅ ✅ ❌ ❌ ❌  → 刚好凑齐 W=4
t4  ❌ ✅ ✅ ✅ ✅ ❌ ❌  → 开始滑动:丢最老的 t0
t5  ❌ ❌ ✅ ✅ ✅ ✅ ❌  → 丢 t1
t6  ❌ ❌ ❌ ✅ ✅ ✅ ✅  → 丢 t2

W=4096 时就是这个模式的放大版,理解它就理解了 SWA 的全部原理。

SWA 落地的四个常见坑(Mistral 用户必看)

坑 1:SWA 原生看不到「文档开头的全局信息」,签约方/标题/摘要类任务直接翻车。 典型案例:128K 长合同最后一段问「合同甲方是谁」,甲方名字只在文档第 5 个 token 出现过,SWA W=4096 的第 127999 个 token 根本看不到第 5 个,直接答「未知」。Fix 有三种(成本递增):(1) 最简单,业务上让用户把问题要的关键信息在 Prompt 结尾重复一遍;(2) 推荐做法:全局锚点注入——把前 256 个 token 的 KV 单独列出来,每个 query 都额外看这 256 个全局 key,不参与滑窗;成本只多 256 个注意力对,几乎零开销;(3) 最稳:分层 Attention——前 2 层用 SWA(局部),最后 2 层用全局注意力(看全部),全局层只开最后两层成本不高,最后一层肯定能看到全局信息。

坑 2:直接把 Llama 2 全注意力权重 + SWA 掩码就跑,代码类任务掉分明显。 原因是代码里经常有「第 5 行定义函数,第 6000 行调用」这种跨 6000+ token 的引用,W=4096 覆盖不到。Fix:(1) 代码场景至少把 W 调到 8192;(2) 用 Rolling Buffer(滚动滑窗)不同层窗口偏移,第 0 层 [0, W]、第 1 层 [W/2, 3W/2]、第 2 层 [W, 2W],相当于不同层的注意力拼起来覆盖更长距离,调用第 6000 行的函数时,某个偏移层的窗口刚好能包含到定义;(3) 做一次「长代码 SFT 适配」:用 100K 条长代码函数调用对继续微调 1 个 epoch,让模型适应 SWA 掩码,代码分基本能涨回来。

坑 3:KV Cache Ring Buffer(环形队列)和 Streaming 流式输出顺序冲突。 全注意力 KV Cache 是追加式写(往数组尾巴加),SWA 是环形写(写到下标 W-1 后跳回 0 覆盖最老),如果推理引擎没处理好,流式输出到第 4097 个 token 时会突然崩或重复吐 token。Fix:不用自己写 Ring Buffer,直接上 vLLM ≥ 0.3.6 或 SGLang,这两家都在 2024 Q2 修完了 SWA KV Cache 的所有边界 bug,开箱即用;不要用 HuggingFace Transformers 的原生 generate 跑 SWA 长上下文,慢而且 bug 多。

坑 4:SWA + RAG 时 chunk 之间的上下文衔接被切断。 典型场景:RAG 把 100K 文档切成 20 个 chunk,每个 chunk 5K 字,分块独立检索后塞进 LLM 上下文做回答,分块边界处的连贯信息(比如段落首句承上启下)被 SWA 看不到。Fix:检索召回的 chunk 之间留 10-15% 的 overlap(相邻 chunk 首尾重叠 500 token),确保任何分块边界的语义都在至少一个 chunk 的 SWA 窗口以内被覆盖。

常见问题

SWA 和 GQA 是什么关系?经常一起出现是因为什么?
它们是正交的两个优化技巧,分别解决 Attention 的两个维度的成本问题——SWA 解决「N×N 矩阵中忽略大部分 query-key 对(行数/列数 N)」,GQA 解决「K/V 头数太多导致 KV Cache 体积大(头数维度)」,两者相乘效果叠加,所以新模型几乎同时开两者(Mistral、Llama 3、Qwen 2.5 全是 SWA+GQA 双开)。类比:Attention FLOPs = (序列长 N) × (头数 H_kv) × (头维 d_k) × (层数 L) × 常数。SWA 把第一个因子 N 从「全部 128K」降到「窗口 W=4K」,缩小 32 倍;GQA 把第二个因子 H_kv 从 32 降到 8 或 4,再缩小 4-8 倍;两者相乘总共省 128-256 倍 KV Cache 体积和 Attention FLOPs,128K 长上下文 7B 模型才能在单张 4090 24G 上跑得动,缺一不可。
如果把 W 设得和最大上下文一样大(比如 N=128K,W=128K),是不是就等价于全注意力?
形式上(掩码上)是的,W=N 时滑窗掩码就是完整的下三角矩阵(causal mask),和全注意力数学上完全等价;但工程实现上(KV Cache 管理方式)可能不一样——很多推理引擎在检测到 W=N 时会自动走 Full Attention 的 KV Cache 追加路径而不是 Ring Buffer 路径,避免环形队列的开销。所以你想在 SWA 模型上临时开全注意力,直接把 sliding_window 参数临时改成等于 max_seq_len 即可,不用重新加载权重,Attention 计算数学等价。注意:如果你训模型时用的是 W=4096 的训练掩码(比如 Mistral 原始训练时只训了 W=4096,推理时开到 32K),那 W=32K 是「训短推长」的 extrapolation(外推),效果会下降(称为 Position Interpolation 问题),必须配合 RoPE 频率压缩 YaRN / NTK-aware Scaling 等位置编码外推技巧才能用;如果模型在预训练阶段就开过大 W(比如 Llama 3.1 预训练就用 W=131072),那推理时开 W=128K 是内插而不是外推,效果完全无损。
我要做 100 万 token 级别的超长上下文(整本书),纯 SWA 够用吗?还是得上其他方案?
纯 SWA 不够,除非你加了多层偏移 + 全局锚点 + 压缩摘要层的多层级 SWA;否则 SWA 最长有效区间在实际生产中大概 64K-128K,1M 级必须上「层次化长上下文方案」混合使用。2025 年百万级上下文成熟的三阶段流水线是:(1)最外层 1M token 级:用「Block-Recurrent Transformer / RWKV / Mamba 等 O(N) 严格线性结构」做压缩状态机,每 4K token 一个块,块之间传一个 2048 维的压缩状态向量,状态向量摘要了前面所有块的精华;(2)中间层 32-64K:用 普通 SWA + 全局锚点(前 256 + 章节开头锚点),对当前要阅读的块做局部精准推理;(3)最外层检索:RAG / Graph RAG 辅助定位,用户问了一个在第 80 万 token 位置的信息,先粗检索找到第 200 个块是相关区域,再把第 198-202 共 5 个块(20K token)切出来送进 SWA 细阅读。别让 LLM 硬吃 1M 全部 token,哪怕 SWA 注意力省了,前面的 FFN 层还是 O(N) 线性,1M token FFN 层计算量是 128K 的 8 倍,P95 延迟还是炸;层次化 + 检索是百万级上下文的唯一正解,SWA 负责最内层精准阅读。