7.3 长上下文与稀疏注意力:滑动窗口、NSA 与 FlashAttention

1.长上下文为什么是块硬骨头

7.2 节我们讲线性注意力怎么把复杂度从 O(n2)O(n^2) 压到 O(n)O(n)。但还有另一条并行的路:不抛弃注意力本身,而是让它变"稀疏"——只算该算的那部分相似度,其余跳过。这条路对那些"必须用真注意力、但又想支持几十万 token 上下文"的场景特别关键。

先说清楚长上下文为什么难。注意力的开销有两块:计算(算 QKQK^\topO(n2)O(n^2))和显存(存注意力矩阵和 KV cache,O(n)O(n) 每层每头,但层数多头数多时总量惊人)。序列拉到 100 万 token 时,光 KV cache 就能吃掉几百 GB 显存,这才是长上下文落地的真正拦路虎。稀疏注意力主要攻计算,KV cache 的显存问题则要靠 9.2 节讲的量化和 8.1 节讲的 PagedAttention 来配合。

2.滑动窗口与局部注意力

最朴素的稀疏化是滑动窗口注意力(sliding window attention):每个位置只和它前后 ww 个位置算注意力,而不是和全部 nn 个位置算。复杂度从 O(n2)O(n^2) 降到 O(nw)O(n\cdot w)。直觉是:大部分时候我们关心的是近距离的上下文(一个词和它前后的词),远处偶尔需要,但不必每个位置都和所有远处位置算。

纯滑动窗口有个问题——信息传不远。第 1 个词要影响第 1000 个词,得经过中间 999 次窗口接力,每次都可能稀释。Mistral 等模型用的招数是让不同层用不同的窗口大小,或者让窗口随层数增大,这样底层看局部、高层看全局,信息能一层层扩散开。这是个简单但有效的工程妥协。

3.可学习的稀疏注意力:NSA 与 DSA

更聪明的稀疏化是让模型自己学该关注哪些位置,而不是人为规定窗口。2025 年 DeepSeek 提出的 NSA(Native Sparse Attention) 是这条线上的代表作,它的核心是三个分支的组合:

  • 压缩分支:把一大段 token 压成一个粗粒度的表示(比如均值池化),用这个粗表示快速判断"这一段值不值得细看"。
  • 选择分支:根据压缩表示打分,挑出最相关的若干个 token 块,只对这些块做细粒度注意力。
  • 滑动窗口分支:保留一个局部窗口,保证近距离细节不丢。

三路并行,再拼起来,模型既看到了全局趋势(压缩),又精准聚焦了关键块(选择),还保留了局部细节(窗口)。关键是整个机制端到端可训练硬件对齐(算子设计成 GPU 友好),所以既能省算力又不掉精度。NSA 的论文(arXiv:2502.11089)是 2025 年长上下文方向被引用最多的之一。

DeepSeek 后续在 V3.2 / V4 上把 NSA 落地成了 DSA(DeepSeek Sparse Attention),用在生产模型里支持超长上下文推理。GLM-5 系列也采用了类似的 DSA 风格稀疏注意力来压低训练和推理成本。这说明可学习的稀疏注意力已经从论文走进了产品

4.FlashAttention:不改变数学,只改变算的方式

有一类工作特别值得一提,它完全不改变注意力的数学,只改变在 GPU 上的计算方式,却能带来数倍加速和大幅显存节省——这就是 FlashAttention 系列(Tri Dao 等人)。

它的核心洞察是:注意力慢,主要慢在反复读写显存(HBM),而不是算本身。传统实现先把整个 QKQK^\top 算出来写回显存,再读出来做 softmax,再写回……GPU 算力其实过剩,瓶颈在显存带宽。FlashAttention 的做法是分块(tiling):把 Q,K,VQ,K,V 切成小块,让每块在 GPU 的快速片上 SRAM 里算完 softmax 再写回,避免中间结果落地到慢速 HBM。数学上结果和标准注意力一模一样,但 IO 大幅减少。

FlashAttention 已经成了 2026 年几乎所有主流大模型的标配——你用的 GPT、Llama、DeepSeek,底下基本都跑着 FlashAttention 或它的变体。它和稀疏注意力是互补的:稀疏注意力减少要算的量,FlashAttention 让剩下的计算跑得更快更省显存。两者叠加,才是今天长上下文能落到工程上的完整答案。

5.把这些和 7.2 串起来

你可能觉得这一节和 7.2 节讲线性注意力有点像,都是在攻 O(n2)O(n^2)。区别要分清:

  • 7.2 的线性注意力/SSM:从根上改变注意力机制,把 softmax 去掉或替换,复杂度降到 O(n)O(n),但牺牲了一点精确检索能力。
  • 本节的稀疏注意力:保留 softmax 注意力的数学,只跳过部分计算,复杂度介于 O(n)O(n)O(n2)O(n^2) 之间,精确性更接近真注意力。

实际工程里,这两种经常配合:用 SSM/线性注意力做主体(兜底效率和长程),在少数关键层用稀疏注意力(保证精确检索),再叠加 FlashAttention 加速——这就是现代长上下文大模型的标准配方。

6.练习

Q1. 滑动窗口注意力把复杂度从 O(n2)O(n^2) 降到多少?它的主要短板是什么,工程上怎么补救?

降到 O(nw)O(n\cdot w)ww 是窗口大小)。短板是信息传不远(远距离依赖要靠多层窗口接力,易稀释)。补救:不同层用不同/递增的窗口大小,让高层看全局、底层看局部,信息逐层扩散。

Q2. NSA 的三个分支各管什么?为什么说它是"端到端可训练"的稀疏注意力?

压缩分支看全局趋势(粗粒度)、选择分支精准聚焦关键块、滑动窗口分支保留局部细节。三者都由可微的网络参数控制,分数和权重都能通过梯度学习,所以整个稀疏模式是模型自己学出来的,而不是人工硬编码。

Q3. FlashAttention 不改变注意力的数学,它是怎么实现加速和省显存的?

瓶颈在显存 IO 而非算力。FlashAttention 用分块(tiling)把 Q,K,VQ,K,V 切小块,让 softmax 在 GPU 快速片上 SRAM 里完成,避免中间大矩阵反复读写慢速 HBM。数学结果不变,IO 大幅减少,从而加速并省显存。

Q4.(大厂面试题) 线性注意力(7.2)和稀疏注意力(本节)都在攻 O(n2)O(n^2),它们的核心区别是什么?为什么生产模型经常把两者混着用?

线性注意力从根本上改机制(去 softmax),复杂度 O(n)O(n) 但精确检索稍弱;稀疏注意力保留 softmax 数学,只跳过部分计算,复杂度居中但精确性更接近真注意力。混用是因为:线性注意力兜底效率和长程,稀疏注意力(在少数层)保证精确检索,再叠 FlashAttention 加速——三者各取所长,是长上下文落地的标准配方。

7.小结

长上下文的工程答案,是把"算得少"(稀疏注意力/NSA)和"算得快"(FlashAttention)叠加起来,再配合线性注意力主体和 KV cache 量化(9.2 节)。但到这里我们只解决了"算力撑不撑得住"这一半——长上下文还有另一半硬骨头:训练时只见过 4k 长度,推理时要撑到 32k,位置编码怎么跟着外推出去。下一篇我们就来啃它,讲清旋转位置编码 RoPE 和它的长度外推变体(NTK、YaRN)。

下一章见。

相关标签
深度学习前沿架构长上下文稀疏注意力FlashAttention