详细解释
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 窗口以内被覆盖。