7.3 长上下文与稀疏注意力:滑动窗口、NSA 与 FlashAttention
1.长上下文为什么是块硬骨头
7.2 节我们讲线性注意力怎么把复杂度从 压到 。但还有另一条并行的路:不抛弃注意力本身,而是让它变"稀疏"——只算该算的那部分相似度,其余跳过。这条路对那些"必须用真注意力、但又想支持几十万 token 上下文"的场景特别关键。
先说清楚长上下文为什么难。注意力的开销有两块:计算(算 ,)和显存(存注意力矩阵和 KV cache, 每层每头,但层数多头数多时总量惊人)。序列拉到 100 万 token 时,光 KV cache 就能吃掉几百 GB 显存,这才是长上下文落地的真正拦路虎。稀疏注意力主要攻计算,KV cache 的显存问题则要靠 9.2 节讲的量化和 8.1 节讲的 PagedAttention 来配合。
2.滑动窗口与局部注意力
最朴素的稀疏化是滑动窗口注意力(sliding window attention):每个位置只和它前后 个位置算注意力,而不是和全部 个位置算。复杂度从 降到 。直觉是:大部分时候我们关心的是近距离的上下文(一个词和它前后的词),远处偶尔需要,但不必每个位置都和所有远处位置算。
纯滑动窗口有个问题——信息传不远。第 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),而不是算本身。传统实现先把整个 算出来写回显存,再读出来做 softmax,再写回……GPU 算力其实过剩,瓶颈在显存带宽。FlashAttention 的做法是分块(tiling):把 切成小块,让每块在 GPU 的快速片上 SRAM 里算完 softmax 再写回,避免中间结果落地到慢速 HBM。数学上结果和标准注意力一模一样,但 IO 大幅减少。
FlashAttention 已经成了 2026 年几乎所有主流大模型的标配——你用的 GPT、Llama、DeepSeek,底下基本都跑着 FlashAttention 或它的变体。它和稀疏注意力是互补的:稀疏注意力减少要算的量,FlashAttention 让剩下的计算跑得更快更省显存。两者叠加,才是今天长上下文能落到工程上的完整答案。
5.把这些和 7.2 串起来
你可能觉得这一节和 7.2 节讲线性注意力有点像,都是在攻 。区别要分清:
- 7.2 的线性注意力/SSM:从根上改变注意力机制,把 softmax 去掉或替换,复杂度降到 ,但牺牲了一点精确检索能力。
- 本节的稀疏注意力:保留 softmax 注意力的数学,只跳过部分计算,复杂度介于 和 之间,精确性更接近真注意力。
实际工程里,这两种经常配合:用 SSM/线性注意力做主体(兜底效率和长程),在少数关键层用稀疏注意力(保证精确检索),再叠加 FlashAttention 加速——这就是现代长上下文大模型的标准配方。
6.练习
Q1. 滑动窗口注意力把复杂度从 降到多少?它的主要短板是什么,工程上怎么补救?
降到 ( 是窗口大小)。短板是信息传不远(远距离依赖要靠多层窗口接力,易稀释)。补救:不同层用不同/递增的窗口大小,让高层看全局、底层看局部,信息逐层扩散。
Q2. NSA 的三个分支各管什么?为什么说它是"端到端可训练"的稀疏注意力?
压缩分支看全局趋势(粗粒度)、选择分支精准聚焦关键块、滑动窗口分支保留局部细节。三者都由可微的网络参数控制,分数和权重都能通过梯度学习,所以整个稀疏模式是模型自己学出来的,而不是人工硬编码。
Q3. FlashAttention 不改变注意力的数学,它是怎么实现加速和省显存的?
瓶颈在显存 IO 而非算力。FlashAttention 用分块(tiling)把 切小块,让 softmax 在 GPU 快速片上 SRAM 里完成,避免中间大矩阵反复读写慢速 HBM。数学结果不变,IO 大幅减少,从而加速并省显存。
Q4.(大厂面试题) 线性注意力(7.2)和稀疏注意力(本节)都在攻 ,它们的核心区别是什么?为什么生产模型经常把两者混着用?
线性注意力从根本上改机制(去 softmax),复杂度 但精确检索稍弱;稀疏注意力保留 softmax 数学,只跳过部分计算,复杂度居中但精确性更接近真注意力。混用是因为:线性注意力兜底效率和长程,稀疏注意力(在少数层)保证精确检索,再叠 FlashAttention 加速——三者各取所长,是长上下文落地的标准配方。
7.小结
长上下文的工程答案,是把"算得少"(稀疏注意力/NSA)和"算得快"(FlashAttention)叠加起来,再配合线性注意力主体和 KV cache 量化(9.2 节)。但到这里我们只解决了"算力撑不撑得住"这一半——长上下文还有另一半硬骨头:训练时只见过 4k 长度,推理时要撑到 32k,位置编码怎么跟着外推出去。下一篇我们就来啃它,讲清旋转位置编码 RoPE 和它的长度外推变体(NTK、YaRN)。
下一章见。