从长文档摘要说起:为什么注意力范围成了瓶颈
假设你正在构建一个企业知识库的摘要服务,需要把一份 50 页的技术白皮书压缩成 500 字的中文摘要。输入序列可能超过 1 万 token。标准 Transformer 的自注意力机制要求每个 token 与序列中所有其他 token 计算关联度,序列长度为 (n) 时,注意力矩阵的元素数量以 (n^2) 的速度增长。当 (n=10000),仅 QK 矩阵乘法就需要约 (10^8) 次浮点运算,KV Cache 的显存占用也随长度线性膨胀。在 16k 上下文的场景下,仅 KV Cache 就可能占用数 GB 显存,普通推理服务器难以承受。
常见的直觉是直接截断输入,只保留最近几千 token。但这会丢失文档前部的关键背景,摘要质量显著下降。另一种思路是使用全局注意力,让每个 token 都能看到所有位置,但这又回到平方复杂度。滑动窗口注意力(Sliding Window Attention,SWA)提供了一种折中:每个 token 只关注固定窗口内的邻居,通过堆叠多层来间接获取远距离信息,同时配合 KV Cache 轮转机制控制显存增长。Mistral 7B 在其技术报告中采用了这一设计,结合分组查询注意力(GQA),在 7B 参数规模下宣称在多个基准上超越 Llama 2 13B。本文以长文档摘要为贯穿场景,拆解 SWA 的机制、实现与失效边界。
注意力计算量的来源:平方项藏在哪
要理解 SWA 节省了什么,先要看清标准注意力中哪些操作随序列长度平方增长。给定序列长度 (s)、模型隐藏维度 (d_{\text{model}})、头数 (n_q)、每头维度 (d_q),一次前向传播中主要操作的浮点运算量(FLOPs)大致如下:
| 操作 | FLOPs(MHA) | 增长阶数 |
|---|---|---|
| QKV 投影 | (6 \times s \times d_{\text{model}}^2) | 线性 |
| QK 矩阵乘法 | (n_q \times 2 \times s^2 \times d_q) | 平方 |
| Softmax | (n_q \times 3 \times s^2) | 平方 |
| 加权求和(应用到 V) | (n_q \times 2 \times s^2 \times d_q) | 平方 |
| 输出投影 | (2 \times s \times d_{\text{model}}^2) | 线性 |
其中 QK 乘法、Softmax 和加权求和三项与 (s^2) 成正比,是长序列下计算开销的主要来源。KV Cache 需要存储每一层的 Key 和 Value,参数量为 (2 \times L \times s \times d_q \times n_{kv})((L) 为层数,(n_{kv}) 为 KV 头数)。若使用半精度浮点数,实际显存再乘以 2。以 Mistral 7B 为例,其配置为 (L=32),(d_{\text{model}}=4096),(d_q=128),(n_{kv}=8),在输入长度 16k 时,KV Cache 约为 2GB。这解释了为什么长序列推理既慢又吃显存。
SWA 的核心机制:局部窗口与跨层感受野
SWA 的思路很直接:把每个 token 的注意力范围限制在一个固定大小的窗口内,窗口大小为 (W),通常只向前看(因果注意力)或前后各看一半。以 Mistral 7B 为例,其窗口大小为 4096,层数为 32。在每一层,token (i) 只能直接关注位置 (i-W+1) 到 (i) 之间的 token。
但直接限制窗口会导致远距离信息完全丢失。SWA 依赖多层堆叠来间接传递信息:第 2 层的 token 可以关注第 1 层窗口内的 token,而这些 token 又关注了更早的位置,因此第 2 层的 token 实际能“看到”第 1 层窗口之外的信息。这种机制与卷积神经网络(CNN)中的感受野类似:3 层 (3\times3) 卷积核的堆叠,能让顶层像素间接看到输入中 (7\times7) 的区域。在 Transformer 中,如果每层窗口大小为 (W),共有 (L) 层,那么理论上最远能传递的距离为 (W \times L)。Mistral 7B 的 (4096 \times 32 = 131072),即约 131k token 的感受野。
这意味着,虽然每个 token 只计算固定窗口内的注意力,但通过层间传递,模型仍然能利用远距离的上下文信息。对于长文档摘要,文档开头的关键定义可以通过多层传递影响到后面的总结位置,而不必在每个位置都直接计算与开头的注意力。
计算与显存的节省:从平方到线性
使用 SWA 后,QK 乘法、Softmax 和加权求和的计算量从 (O(s^2)) 降为 (O(s \times W))。当 (W) 固定(如 4096),序列长度 (s) 从 4k 增长到 131k 时,这三个操作的 FLOPs 只随 (s) 线性增长,而不是平方增长。Mistral 7B 的技术报告提到,在序列长度 16k、窗口大小 4096 时,将 SWA 实现在 FlashAttention 和 xFormers 中,相比标准注意力获得了约 2 倍的速度提升。
KV Cache 的节省更为直接。标准注意力中,KV Cache 随序列长度线性增长,每个 token 的 Key 和 Value 都需要缓存。而 SWA 中,超出窗口范围的 token 不再需要缓存,因此 KV Cache 的最大大小被限制为 (2 \times L \times W \times d_q \times n_{kv}),与序列长度无关。当序列长度远大于窗口时,显存占用大幅降低。例如,131k 序列长度下,KV Cache 理论最大可节省 (31/32) 的显存(因为窗口只占序列的 (1/32))。
这种有界缓存还带来了工程上的好处:推理时 KV Cache 大小可预先分配,不会因个别超长请求导致显存溢出,便于估算吞吐量。
与 KV Cache 的配合:轮转替换策略
在自回归推理中,标准注意力会缓存所有历史 token 的 Key 和 Value,且只增不减。SWA 则采用轮转替换策略:当新 token 进入窗口时,最旧的 token 的 KV 被丢弃,缓存大小保持为窗口大小。例如窗口 (W=4),处理第 5 个 token 时,第 1 个 token 的 KV 被移除,缓存中始终只有最近 4 个 token 的 KV。
这种策略的实现需要特殊的数据结构。常见做法是使用环形缓冲区(circular buffer),索引取模即可定位。在长文档摘要场景中,输入是一整篇文档,但推理时仍逐 token 生成。生成第一个 token 时,需要处理整个输入序列,此时 SWA 的窗口会覆盖所有输入 token(如果输入长度小于窗口),但 KV Cache 只保留窗口内的部分。随着生成的进行,缓存不断轮转,丢弃最早的输入 token 的 KV。
需要注意,轮转替换只影响 KV Cache,不影响模型对远距离信息的利用——因为跨层传递已经在前向计算中完成。但这也意味着,如果文档超过窗口大小,模型在生成时无法直接“回看”窗口外的原始 token,只能依赖高层已经编码的语义信息。
长 Prompt 的分块处理与工程实现
实际应用中,RAG 或工具调用常常使用固定的长 Prompt。当 Prompt 长度超过窗口大小时,需要分块处理。一种做法是将长 Prompt 切分为多个块,每个块单独计算注意力,块与块之间通过全局 token 或跨层传递连接。Mistral 的官方实现中,长输入会被自动分块,每块大小不超过窗口大小,然后按顺序处理。
工程实现上,SWA 通常通过修改注意力掩码(mask)来实现。标准因果注意力的掩码是下三角矩阵,SWA 则在此基础上增加窗口限制,形成带状掩码。FlashAttention 和 xFormers 等高效注意力库已支持自定义掩码,因此可以复用其优化内核。Mistral 7B 的官方代码中,SWA 通过设置 sliding_window 参数启用,并配合 GQA 使用。
窗口大小与困惑度的权衡
窗口大小 (W) 是 SWA 的核心超参数,直接影响模型质量。窗口过小,模型无法捕捉足够的局部上下文,远距离信息传递效率降低,困惑度(perplexity)上升。窗口过大,计算量和显存占用增加,接近标准注意力。Mistral 7B 选择 4096,是基于语言建模任务中局部依赖占主导的经验。
理论上,窗口大小与困惑度之间存在权衡。较小的窗口能显著降低计算成本,但可能损害对长距离依赖的建模能力。例如,在长文档摘要中,如果窗口只有 512,模型可能无法通过 32 层传递覆盖文档开头,导致摘要遗漏关键信息。相反,窗口设为 8192 可能提升质量,但计算量翻倍。
资料中并未提供 Mistral 7B 在不同窗口下的困惑度对比数据,但可以推断,窗口大小应至少覆盖常见的局部依赖长度,同时考虑层数,使感受野覆盖目标序列长度。实际应用中,通常通过验证集困惑度来调优。
与全局注意力、稀疏注意力的对比
SWA 并非唯一的长序列注意力方案。Longformer 提出了结合局部窗口与全局注意力的稀疏注意力:除局部窗口外,还选择部分 token(如 [CLS] 或特定位置)作为全局 token,它们能关注整个序列,其他 token 也能关注它们。这种设计在保持线性复杂度的同时,提供了显式的全局信息通路。
下表对比了三种方案在长文档摘要场景下的关键特性:
| 方案 | 计算复杂度 | KV Cache 上限 | 全局信息获取方式 | 实现复杂度 | 适用场景 |
|---|---|---|---|---|---|
| 全局注意力 | (O(n^2)) | 随 (n) 增长 | 直接全连接 | 低 | 短序列(<2k) |
| 滑动窗口注意力 | (O(n \times W)) | (O(W)) | 跨层传递 | 低 | 长文档、流式生成 |
| 稀疏注意力(Longformer 式) | (O(n \times (W+S))) | (O(W+S)) | 显式全局 token | 中 | 需要显式全局信息的任务 |
SWA 的优势在于实现简单、显存有界、计算效率高,适合对延迟敏感的在线服务。其劣势在于远距离信息依赖层数传递,如果层数不足或窗口过小,可能丢失关键全局信息。稀疏注意力通过添加全局 token 缓解了这一问题,但增加了实现复杂度和计算开销。
在长文档摘要场景中,如果文档结构清晰(如章节标题),稀疏注意力能直接关注标题,效果可能更好;如果文档是连续叙述,SWA 的跨层传递通常足够。
失败模式与可观测性
SWA 并非万能,存在一些典型失败模式。
分布偏移:SWA 在预训练时窗口大小固定,如果下游任务输入长度远超预训练长度,模型可能无法有效利用远距离信息,即使感受野足够。例如,Mistral 7B 预训练时使用 8k 序列,但窗口 4096 与层数 32 给出的理论感受野为 131k,实际效果可能打折。
长尾输入:当输入中存在与当前 token 相关但距离超过窗口的 token 时,模型无法直接关注,只能依赖高层语义。如果高层未能充分编码该信息,输出可能不准确。
缓存轮转的副作用:轮转替换会丢弃早期 token 的 KV,如果后续生成需要重新参考这些 token(如摘要中需要引用文档开头的术语),模型只能依赖已编码的表示,可能丢失细节。
可观测性:生产环境中应监控以下指标:KV Cache 命中率(实际使用的缓存比例)、平均生成延迟、显存占用、困惑度或下游任务指标。当困惑度异常升高时,应检查输入长度是否接近或超过窗口大小,以及是否存在长距离依赖被截断的情况。
部署边界与未解决问题
SWA 适用于需要处理超长序列且对延迟和显存敏感的场景,如长文档摘要、对话系统、代码生成。但它不适用于需要精确引用远距离事实的任务,除非结合检索或全局注意力。
当前 SWA 的一个未解决问题是窗口大小的自适应调整。固定窗口无法适应不同任务对局部和全局依赖的不同需求。一些研究尝试学习注意力模式,但尚未成为主流。另外,SWA 与位置编码的交互也需注意:窗口外的 token 位置信息无法直接获取,可能影响相对位置编码的效果。
在实际部署中,建议先用小规模数据测试不同窗口大小对下游任务的影响,再根据延迟和显存预算选择。同时,结合 GQA 或 MQA 可进一步降低 KV Cache 大小,但会牺牲少量质量。
贯穿场景的流程图
以下流程图展示了长文档摘要服务中,SWA 从输入到输出的完整流程:
flowchart TD
A[输入长文档] --> B[分块处理]
B --> C[逐块计算 SWA 注意力]
C --> D[更新 KV Cache 轮转替换]
D --> E{是否生成结束?}
E -- 否 --> F[生成下一个 token]
F --> C
E -- 是 --> G[输出摘要]
C --> H[跨层传递远距离信息]
H --> C
在分块处理阶段,长文档被切分为不超过窗口大小的块。每个块内计算 SWA 注意力,同时通过层间传递获取块外信息。KV Cache 在块间轮转,只保留当前窗口内的 KV。生成阶段,每个新 token 的注意力计算只涉及窗口内的 KV,因此延迟稳定。
结论
滑动窗口注意力通过限制注意力范围,将计算复杂度从平方降为线性,并通过跨层传递和 KV Cache 轮转保留远距离信息。它在长文档摘要等场景中提供了显存可控、延迟稳定的推理能力,但窗口大小的选择、远距离依赖的可靠性仍是需要权衡的工程问题。理解其机制与边界,有助于在具体场景中做出合理的架构决策。