DeepSeek NSA:Sparse Attention 不是新概念,但这次真的做对了
DeepSeek NSA:Sparse Attention 不是新概念,但这次真的做对了
原文:Native Sparse Attention: Hardware-Aligned and Natively Trainable Sparse Attention
1. 前言:先把背景交代清楚
你有没有想过这样一个问题:为什么模型上下文越长,生成越慢、服务器账单越贵?
Attention 机制在这里有两笔躲不掉的账,都跟序列长度强相关:
- 计算量:理论上限是 $O(N^2)$,序列翻倍,prefill 的计算量翻四倍;
- 内存访问:每一个历史 token 的 K/V 都要保存,decode 阶段每生成一个新 token,都得把过去所有的 K/V 从 HBM 里读一遍——而 decode 恰恰是 memory-bound(内存带宽受限)的,读多少数据,就要等多久。
所以上下文能力是”真香”的,attention 成本也”真贵”,64k、128k 这样的长上下文场景里,attention 的账单几乎肉眼可见地涨。大家的第一个念头很自然:反正也不是每个历史 token 都有用,跳过不重要的,是不是就能便宜下来?
这个想法有个正式名字:Sparse Attention(稀疏注意力)。而且它一点都不新——BigBird、Longformer、Reformer,2020 年那波论文基本把”哪些 token 可以不算”这件事在算法层面研究透了:局部窗口、块状稀疏、全局哨兵 token,该有的花样都有了。
但真正好用的 sparse attention 工程实现,现在才开始出现。
这不是贬低前人,是实事求是:以前的 sparse attention 在 GPU 上的落地大多很尴尬——计算上确实跳过了部分 token,但内存访问 pattern 还是按 dense 来的,HBM 带宽该占多少占多少,SRAM 利用率也差。结果就是理论 FLOPs 降了不少,实际 wall-clock time 没快多少。而且更麻烦的是可训练性:很多方案是拿 full attention 训练完的模型,推理时再做近似(H2O、StreamingLLM 都走这条路线),模型从来没为稀疏模式优化过,近似误差是固定的,质量必然有折扣。
大模型时代的 attention 优化有个公认的分水岭——FlashAttention 证明了一件事:attention 慢的根子不在 FLOPs,在 HBM 访问;搞定内存访问比省计算量更重要。 DeepSeek 这篇 NSA(Native Sparse Attention)做的事,大致可以理解成”把 FlashAttention 的这条路在稀疏场景下重新走一遍”:
算法设计和 GPU kernel 实现一起考虑,让 sparsity 真正落地为实际加速;同时支持直接从零预训练稀疏模型,而不是事后近似。
64k 长度序列上,对比 FlashAttention-2:decode 快 11.6×,forward 快 9×,backward 快 6×,而且模型质量不降反升——通用 benchmark 平均分 0.456 vs Full Attention 的 0.443,LongBench 上 0.469 vs 0.437,长推理链的 AIME 上优势更大。一张图看全貌:

这三个数字里 decode 反而提升最大,原因值得先说一下,它也是理解这篇论文的关键:
- Decode:每生成一个 token,都把历史 K/V 从 HBM 读一遍,是 memory-bandwidth bound。稀疏化直接跳过大部分 KV 的 load,HBM 读取量大幅下降,对带宽受限的操作收益最大,所以反而最快——11.6×
- Forward(prefill):整个 prompt 并行处理,是 compute-bound,省下的 FLOPs 直接转化为加速——9×
- Backward:还要算梯度,attention 之外的大量操作没法稀疏,折扣最大——6×
好,背景交代完毕。下面先把”老 sparse attention 为什么没做对”拆开讲,再看 NSA 的算法和 kernel 是怎么对症下药的。
2. 老 Sparse Attention 为什么没做好:算法对,硬件不认
先说清楚问题在哪,这样才能理解 NSA 解决了什么。
Attention 的计算复杂度是 $O(N^2)$,序列长度翻倍,计算量翻四倍。Sparse attention 的思路是:不是所有 token 对之间的 attention 都有意义,跳过那些”不重要的”token,把 $O(N^2)$ 压下来。
算法层面没问题。但落到 GPU 上:
GPU 喜欢连续、规整的内存访问。 你告诉它”第 17、43、128、512 号 token 的 KV 我要,其他不要”,GPU 就得在 HBM 里东一块西一块地 load,内存访问变成 random access,带宽利用率极差。哪怕你 FLOPs 少了 90%,memory bound 没解决,速度未必快多少。
另一个问题是可训练性。 很多 sparse attention 是训练完一个 full attention 模型,推理时再做近似(比如 H2O、StreamingLLM)。这种方式不可能做到最优——模型本身没有针对 sparse pattern 训练优化过,近似误差是固定存在的。
NSA 把这两个问题一起解决。
3. NSA 的三路并行架构:粗读、精读、近读
算法层面,NSA 把 attention 拆成三个并行分支,每个分支处理不同粒度的信息,最后用 gate 加权合并。说人话就是三种阅读策略各司其职:全局走势靠压缩总结(粗读)、关键位置靠精读原文(精读)、近处靠逐词看(近读):

先把符号定好,后面代码和维度全靠它。以论文的 27B 配置为例:
- 序列长 $t$,64 个 query head,GQA 分 4 组(每组 16 个 head 共享一份 KV),$d_k = 192$,$d_v = 128$
- 压缩分支:块长 $l = 32$,步长 $d = 16$
- 选择分支:块长 $l’ = 64$,每 query 选 $n = 16$ 个块(其中 1 个初始块 + 2 个局部块固定激活)
- 滑窗分支:窗口 $w = 512$
3.1 Compressed Attention(粗读):把历史压成几百个 token
这一路把整个 KV 历史按块压缩——每 $l$ 个 token 用一个 MLP(带块内位置编码)压成一个代表向量,然后对这些压缩后的 KV 做标准 attention:
\[\tilde{\mathbf{K}}_t^{\text{cmp}} = \{\,\varphi(\mathbf{k}_{id+1:id+l}) \mid 0 \le i \le \lfloor (t-l)/d \rfloor \,\}\]$\varphi$ 就是那个压缩 MLP。注意步长 $d < l$,相邻块有重叠——避免关键信息正好被块边界切成两半。压缩后 token 数约 $t/d$,8k 序列只剩约 511 个,attention 量直接降一个量级。
import torch
import torch.nn as nn
import torch.nn.functional as F
class CompressionBranch(nn.Module):
"""粗读:把 KV 历史压成 ~t/d 个 token 再做 attention"""
def __init__(self, d_k=192, d_v=128, l=32, d=16):
super().__init__()
self.l, self.d = l, d
# φ:块内 l 个 token(展平后 l*d_k 维)→ 一个压缩向量
self.phi_k = nn.Linear(l * d_k, d_k, bias=False)
self.phi_v = nn.Linear(l * d_v, d_v, bias=False)
def forward(self, q, k, v):
# q: [H, d_k](一个 GQA 组内的 query heads)
# k: [t, d_k],v: [t, d_v](本组共享的 KV 历史)
blocks_k = k.unfold(0, self.l, self.d) # [n_cmp, d_k, l],滑窗切块,d < l 有重叠
blocks_v = v.unfold(0, self.l, self.d) # [n_cmp, d_v, l]
n_cmp = blocks_k.shape[0] # 压缩 token 数:t=8192 → 511
k_cmp = self.phi_k(blocks_k.permute(0, 2, 1).flatten(1)) # [n_cmp, d_k]
v_cmp = self.phi_v(blocks_v.permute(0, 2, 1).flatten(1)) # [n_cmp, d_v]
p_cmp = torch.softmax(q @ k_cmp.transpose(-1, -2) / q.shape[-1] ** 0.5, dim=-1)
# ↑ [H, n_cmp] 压缩注意力分布——注意这个要留着,选择分支直接复用
return p_cmp @ v_cmp, p_cmp # 输出 [H, d_v] + 注意力分布
保留全局上下文的粗粒度感知是这一路的核心作用:不管序列多长,模型都能”看到”整个历史的大致走向,不会因为稀疏化丢掉远距离信息。
3.2 Selected Attention(精读):top-n 块选择
这一路做精细选择:挑出最重要的 $n$ 个块,对选中块的原始 KV(不压缩)做完整 attention。比如写代码时当前 token 要精确 attend 到几百行之前的变量定义——粗读负责定位”大概在哪”,精读负责精确计算。
第一个问题:怎么知道哪些块重要? 单独训一个重要性预测网络?NSA 的答案是白嫖:直接复用粗读分支的注意力分数 $p_t^{\text{cmp}} = \text{Softmax}(q_t^\top \tilde{\mathbf{K}}_t^{\text{cmp}})$——压缩块和原始块位置对齐,压缩分数天然就是块重要性估计,零额外计算。这个设计后面(第 5 节)你会看到它有多关键。
第二个问题:为什么按”块”选,而不是按 token 选? 除了内存连续性,还有个注意力本身的性质——论文可视化了训练好的 27B full attention 模型的 attention map:

attention 分数呈明显的块状聚集分布:相邻的 key 位置分数高度相关(这个现象 MInference 也观察过)。既然重要的 token 本来就成片出现,按块选几乎不损失召回率。
选择的具体流程(对应论文 Eq.8-12):
class SelectionBranch(nn.Module):
"""精读:按压缩分数挑 top-n 个原始 KV 块,做精确 attention"""
def __init__(self, l_prime=64, n=16, d=16):
super().__init__()
self.l_prime, self.n, self.d = l_prime, n, d
def forward(self, q, k, v, p_cmp):
# q: [H, d_k],k: [t, d_k],v: [t, d_v]
# p_cmp: [H, n_cmp],粗读分支白送的块重要性分数
t, H = k.shape[0], q.shape[0]
n_blk = t // self.l_prime # 选择块数:t=8192, l'=64 → 128 块
# ① 块分数对齐:一个选择块(l'=64)覆盖 l'/d=4 个压缩步长,
# 把落在它里面的压缩分数求和聚成块分数
# (真实实现对齐边界有专门处理,这里 pad 齐再 reshape,教学简化)
n_per = self.l_prime // self.d # 每个选择块覆盖的压缩步长数 = 4
p_pad = F.pad(p_cmp, (0, n_blk * n_per - p_cmp.shape[-1])) # [H, 512]
p_blk = p_pad.view(H, n_blk, n_per).sum(-1) # [H, n_blk=128]
# ② GQA 组内求和:同组 head 共享 KV,必须选同一批块,否则 KV load 无法复用
p_grp = p_blk.sum(0) # [n_blk]
# ③ top-n 选块(离散操作,梯度不穿过这里,见第 5 节)
idx = p_grp.topk(self.n).indices # [n],例:[3, 17, 42, ...]
# ④ 按索引取连续块——kernel 里唯一的"索引读",但粒度是整块,内存仍然连续
k_sel = k.view(n_blk, self.l_prime, -1)[idx].flatten(0, 1) # [n·l', d_k] = [1024, d_k]
v_sel = v.view(n_blk, self.l_prime, -1)[idx].flatten(0, 1) # [n·l', d_v]
# ⑤ 组内每个 head 对同一批选中 KV 做精确 attention
return F.scaled_dot_product_attention(q, k_sel, v_sel) # [H, d_v]
第 ② 步值得单独强调:GQA 组内所有 query head 必须选同一批块。如果每个 head 各选各的,共享的 KV 就要按所有 head 选择的并集来 load,GQA 省带宽的优势直接作废——这是论文里点名批评 Quest 类方法的一点。
3.3 Sliding Window(近读):最简单也最必要
这一路对最近 $w = 512$ 个 token 做标准 attention,FlashAttention 直接支持。语言本质上是局部依赖的,相邻 token 关联最强,这一路保证局部连贯性不被前两路遗漏。
但它还有一个容易被忽略的作用:防止”捷径学习”。局部模式是最容易学的,如果只有压缩和选择两路,模型可能偷懒只学局部,长程建模能力得不到训练。单独划一条滑窗支路把局部需求隔离掉,粗读和精读才能被迫去学长距离模式。
论文还有个配套设计:三路各自有独立的 K/V 投影矩阵,不共享。理由同上——共享 KV 的话梯度会互相干扰,局部模式照样会劫持学习信号。
3.4 Gate 合并与总装:一次 attention 调用的完整数据流
三路输出用可学习的 gate 加权合并:
\[\mathbf{o}_t^* = \sum_{c \in \{\text{cmp, slc, win}\}} g_t^c \cdot \text{Attn}(\mathbf{q}_t, \tilde{\mathbf{K}}_t^c, \tilde{\mathbf{V}}_t^c)\]$g_t^c \in [0,1]$ 由输入特征过一个 MLP + Sigmoid 得到——每个 token 自己动态决定三路怎么分配权重,而不是一个固定的超参。近读分支最简单,先补上它的定义,再把三路总装起来:
class WindowBranch(nn.Module):
"""近读:只看最近 w 个 token,标准 attention,FlashAttention 直接支持"""
def __init__(self, w=512):
super().__init__()
self.w = w
def forward(self, q, k, v):
return F.scaled_dot_product_attention(q, k[-self.w:], v[-self.w:]) # [H, d_v]
class NSAttention(nn.Module):
"""NSA 总装:三路并行 + gate 加权合并(省略 batch 维,单 GQA 组视角)"""
def __init__(self, d_model=2560, n_heads=16, d_k=192, d_v=128):
super().__init__()
self.wq = nn.Linear(d_model, n_heads * d_k, bias=False) # [d_model → 16×192]
# 三路独立 K/V 投影(防 shortcut)
self.wk = nn.ModuleDict({c: nn.Linear(d_model, d_k, bias=False) for c in ['cmp', 'slc', 'win']})
self.wv = nn.ModuleDict({c: nn.Linear(d_model, d_v, bias=False) for c in ['cmp', 'slc', 'win']})
self.gate = nn.Linear(d_model, 3) # g_t = Sigmoid(MLP(x_t))
self.cmp, self.slc, self.win = CompressionBranch(), SelectionBranch(), WindowBranch()
self.wo = nn.Linear(n_heads * d_v, d_model, bias=False)
def forward(self, x): # x: [t, d_model],t=8192 为例
t = x.shape[0]
q = self.wq(x).view(t, 16, 192) # [t, H=16, d_k]
o_c, p = self.cmp(q, self.wk['cmp'](x), self.wv['cmp'](x)) # 粗读: [16, 128] + [16, 511]
o_s = self.slc(q, self.wk['slc'](x), self.wv['slc'](x), p) # 精读: [16, 128],复用 p
o_w = self.win(q, self.wk['win'](x), self.wv['win'](x)) # 近读: [16, 128]
g = torch.sigmoid(self.gate(x)) # [t, 3],逐 token 的三路权重
o = g[..., 0:1] * o_c + g[..., 1:2] * o_s + g[..., 2:3] * o_w # [t, 16, 128]
return self.wo(o.flatten(-2)) # [t, d_model]
3.5 算一笔维度账:为什么 decode 是 11.6×
把三路在一个 query 眼里”实际触碰的 KV token 数”加一下,论文的加速比直接就出来了:
| 分支 | 8k 序列 | 64k 序列 |
|---|---|---|
| 粗读(压缩 token) | $(8192-32)/16+1 = 511$ | $(65536-32)/16+1 = 4096$ |
| 精读($n \cdot l’$) | $16 \times 64 = 1024$ | $16 \times 64 = 1024$ |
| 近读($w$) | $512$ | $512$ |
| 合计 | ≈2047 | ≈5632 |
| Full Attention | 8192 | 65536 |
| 理论加速比 | 4× | 11.6× |
和论文 Table 4 的数字严丝合缝:decode 是 memory-bound,省多少 KV load 就快多少倍。注意精读和近读的触碰量是常数(不随 $t$ 增长),只有粗读随 $t/d$ 线性涨——这就是”序列越长加速比越高”的来源。
4. Hardware-Aware Kernel 设计:把稀疏变成”可执行的快”
三路架构是 NSA 能快起来的算法基础,但真正的加速来自 kernel 实现。如下图,NSA 的 Triton kernel 设计图:

两个关键设计:
1. GQA 分组共享 KV load(Grid Loop)
FlashAttention 按时间连续块 load query,但同一个 query 块内的各个 head 可能需要完全不同的 KV 块,SRAM 利用率很差。NSA 换了个分组策略:按 GQA 组组织 query——对每个位置,把同组所有 16 个 query head 一起 load 进 SRAM,因为它们共享同一批稀疏 KV 块,KV 只需 load 一次整组复用。
结果:HBM → SRAM 的 load 次数大幅减少,带宽利用率显著提升。
2. 块级 KV 选择(Inner Loop)
Selected attention 的稀疏索引是按块存的连续 index,Inner Loop 按这个 index 顺序 load 对应 KV 块进 SRAM 再计算。因为选中的是整块,load 全是顺序 coalesced access,不存在 random access 的问题。
这和”每次 load 一个 token 的 KV”有本质区别——前者是顺序化的结构,后者是跳变的随机。
顺带一提,三路里压缩和滑窗两路都是规整的 attention,FlashAttention-2 原版 kernel 直接能用,只有精读分支需要专门写 kernel——工程量被算法设计压到了最小。
5. 端到端可训练:离散选择的梯度,是这篇论文最值得学的取舍
NSA 支持从头预训练,这一点值得单独展开。三路里有两路天生可微——粗读和近读就是标准 attention,套了个压缩/截断而已;麻烦的只有精读的 top-n 块选择:离散操作,梯度断路。
业界处理这个问题有两条常规路线,DeepSeek 把两条都试了一遍,然后全放弃了:
- 辅助损失路线(SeerAttention 式):额外加一套参数预测块重要性,用真实 attention 分数的块内均值做 KL 监督。问题是:额外算子开销 + 辅助损失经常损害模型性能
- 免参数启发式路线(Quest 式):用 query 和块内 key 的 min-max 统计量的乘积选块。问题是:召回率低;论文还试了”先 1000 步 full attention 冷启动再切换”,也没救回来
3B 模型上的对比非常直观——两条替代路线的 loss 全程压不过 NSA,甚至压不过 full attention:

NSA 的答案是把问题绕过去:块重要性直接复用粗读分支的注意力分数,不加任何参数、不加任何辅助损失。梯度路径是这样的:
- top-n 索引本身不可微,梯度不穿过它——但被选中块的 attention 计算完全可微,梯度正常回传到三路的 Q/K/V 投影和压缩 MLP
- 关键在于”学哪里重要”这个任务被转移给了粗读分支:粗读的压缩分数在语言建模损失下端到端训练,压缩学得越准,选块就越准——两个任务共享一套梯度信号,不需要额外的监督
配套的还有两个训练稳定性设计:
- 三路独立 K/V 投影(3.3 节提过):本质是给梯度通路做隔离,防止局部模式的捷径学习劫持长程分支的训练
- MoE 第一层换成 SwiGLU:backbone 层面的稳定性处理
效果如何?27B 总参数(3B active,MoE)、260B token 预训练,NSA 的 loss 全程压在 Full Attention 下面——稀疏化之后训练 loss 反而更低:

下游更能说明问题。用 R1 蒸馏数据做 SFT 后考 AIME 24(数学推理,长推理链场景):
| 生成长度限制 | 8192 | 16384 |
|---|---|---|
| Full Attention-R | 0.046 | 0.092 |
| NSA-R | 0.121 | 0.146 |
8k 限制下 NSA-R 是 Full Attention-R 的 2.6 倍。推理任务恰恰最依赖长程精确检索(推着推着要回去查前面的中间结论),这正是”原生稀疏训练出来的模型更适配自己的稀疏 pattern”的最好证据——事后近似的方案连训练都不支持,根本没资格上这个牌桌。
6. 实验结果
6.1 通用 Benchmark
如下表,NSA 在大多数通用 benchmark 上超过了同等规模的 Full Attention baseline:平均分 0.456 vs Full Attention 的 0.443,9 项里赢了 7 项(DROP +0.042、GSM8K +0.034):

这是在整个模型计算量大幅下探的情况下做到的。核心原因:稀疏注意力过滤了大量无意义 token 对的干扰,某种程度上相当于对 attention pattern 做了一次”降噪”。
LongBench 上和同类稀疏方案的对比更残酷(所有稀疏方法统一 2560 token 激活预算,含 128 sink + 512 局部):H2O 0.303、InfLLM 0.383、Quest 0.392、Exact-Top(用真实 attention 分数选块的 oracle)0.423、Full Attention 0.437,NSA 0.469——比 oracle 还高 0.046,比 Full Attention 高 0.032。多文档 QA(HPQ +0.087)和代码补全(LCC)提升最大,恰好是”远距离精确检索”型任务。
6.2 长上下文:Needle-in-a-Haystack
64k 上下文的 Needle-in-a-Haystack,NSA 全绿:所有位置、所有深度 accuracy = 1.0,没有任何折损。说明 Compressed + Selected 的双路设计确实兜住了长距离依赖,没有因为稀疏而丢关键信息:

6.3 实际加速(对比 Triton 版 FlashAttention-2)
如下图:

- 8k 序列:forward 2.1×,backward 1.1×(短序列时 sparsity 收益还不明显)
- 32k 序列:forward 6.3×,backward 3.4×
- 64k 序列:forward 9.0×,backward 6.0×
decode 的账在 3.5 节已经算过了,实测和理论对得上:8k → 4×、16k → 6.4×、32k → 9.1×、64k → 11.6×。序列越长,收益越大——这正是 sparse attention 应该有的表现:$O(N^2)$ vs $O(N \cdot k)$ 的差距要序列够长才放得出来。
7. 和以往 Sparse Attention 的本质区别
| 维度 | 老方案(H2O/StreamingLLM 等) | NSA |
|---|---|---|
| 训练方式 | 训练 Full Attention,推理时近似 | 原生稀疏预训练 |
| sparsity 粒度 | token 级(random access) | block 级(coalesced access) |
| 实际加速 | 理论 FLOPs 降,wall-clock 未必 | 显著真实加速(6-11.6×) |
| 质量 | 通常有精度损失 | 达到甚至超过 Full Attention |
最本质的区别就一句话:以前是把 dense 模型想办法变稀疏,NSA 是把稀疏作为训练目标本身。
8. 个人 Take
这篇工作让我想起 FlashAttention 刚出来时的感觉:大家都知道 attention 慢,也都在想各种算法层面的优化,但 FlashAttention 的核心洞察是”IO-awareness 比 FLOPs 更重要“——用 tiling 和 online softmax 在不改变结果的前提下大幅减少 HBM 访问,才真正快了。
NSA 有两层值得学的:
第一层是老话题的新解法:块级选择把 random access 变成 sequential access,配合 GQA 组内共享 load,让 GPU 真正高效执行稀疏计算。
第二层是我认为更通用的设计哲学:当某个操作不可微时,不要急着给它造梯度(辅助损失、straight-through 这些),先看看能不能让系统里已有的可微部分顺带把这个决策学了。 NSA 的 top-n 选择不可微,DeepSeek 没有给它加任何监督,而是让压缩分支的注意力分数兼任选块信号——一套参数干两件事,训练信号干净,实验证明比专职的重要性预测网络更好。这种”借道”的思路在很多离散决策场景(路由、剪枝、量化位分配)都值得想一想。
如果这篇文章涉及的 LLM 推理效率优化你想系统深入,可以看看我之前出版《动手学 AutoML:从 NAS 到大语言模型优化实战》,书里有专章讲 LLM 推理效率和 attention 优化工程背景,正好和本文的场景呼应。
