NSA 原生稀疏注意力(二):Triton Kernel 设计与 27B 模型实验

9021 字
45 分钟
NSA 原生稀疏注意力(二):Triton Kernel 设计与 27B 模型实验

回顾与路线图#

上一篇讲完了 NSA(Native Sparse Attention,原生稀疏注意力)的算法骨架:用「压缩-选择-滑窗」三分支替代全注意力,压缩分支的 softmax 分数同时充当选择分支的块寻址信号,top-n 块选择因此退化为纯内存寻址操作,梯度沿着可微的压缩分支回传,实现了端到端原生训练;按论文配置,64k 上下文下每个 query 只需处理 5632 个键,约为全注意力的 1/11.6。

但「键数降到 1/11.6」和「解码延迟真的降 11.6 倍」之间还隔着一整个工程层:GPU 只认得连续的内存访问和分块的矩阵乘,稀疏算法给出的「随机散布的块索引」如果不经过精心组织,理论上的计算节省会几乎全部浪费在访存上。本文回答三个问题:NSA 的 kernel 是怎么把键数上的节省变成真实加速的(实测前向 9×、反向 6×、解码 11.6×);27B 模型大规模预训练的真实成绩单如何;以及这套思想如何在 DeepSeek-V3.2 的 DSA(DeepSeek Sparse Attention,DeepSeek 稀疏注意力)中完成工程化落地。

从算法到 kernel:理论稀疏与真实加速之间隔着三道坎#

先回到第一篇定义过的算术强度(arithmetic intensity):计算量与访存量的比值,决定任务是算力受限(compute-bound)还是带宽受限(memory-bound)。稀疏注意力要真正变快,必须在两个阶段分别兑现:训练和预填充阶段要省计算(让矩阵乘变小),解码阶段要省访存(让 HBM 流量变小)。但这只是必要条件。即使键数下降 11.6 倍,把「选中的稀疏块」真正在 GPU 上算出来仍然有三个具体障碍,这正是 NSA 论文 kernel 设计的出发点(论文 3.4 节)。

障碍一:FlashAttention 式加载策略与稀疏 KV 集合根本不兼容。 FlashAttention 系 kernel 的工作单元是「连续的一段 query 块 × 全序列 KV」:把一段连续的 query token(比如 64 或 128 个)连同全部历史 KV 分块循环处理,KV 块按序列顺序逐个流经 SRAM。这个策略假设「同一个 query 块内的 token 面对的是同一段连续的 KV」。稀疏选择打破了这个假设——同一个 query 块里的不同 token 各自选中的 KV 块集合互不相干(论文原文用 disjoint 描述),如果照搬 FA 的加载策略,每个 query token 都要抓取自己专属的、散落在缓存各处的 KV 块,内存访问变成 gather 式的随机读,带宽利用率断崖式下跌。论文明确指出:「如果我们遵循 FlashAttention 将时间上连续的 query 块加载进 SRAM 的策略,会导致低效的内存访问,因为块内的 query 可能需要不相交的 KV 块。」

障碍二:GQA 组内共享与逐头选择冲突。 现代模型用 GQA(Grouped-Query Attention,分组查询注意力)让组内多个 query 头共享同一份 KV 缓存,解码时 KV 只加载一次、组内所有头复用。第一篇已经讲过算法层的含义(组内分数求和、选择结果共享);kernel 层的问题更直接:如果每个 query 头各自抓取自己选中的 KV 块,共享的 KV 缓存实际上被加载了头数次,带宽开销乘以头数,GQA 的收益被吃光。KV 传输必须在「组」的粒度上合并。

障碍三:SM 间负载不均衡。 一个 kernel 内部的循环如果长度取决于每个 query 的实际选择结果,不同流式多处理器(streaming multiprocessor,SM)拿到的活就可能差很多——有的 SM 循环 16 个块,有的循环 3 个块,快的 SM 干等慢的,整个 GPU 的空闲时间由最慢的那个 SM 决定(俗称 tail effect)。要让所有 SM 均衡,需要把「工作量几乎相同」的计算单元平铺到调度器上。

NSA 的 kernel 设计(本节描述的设计于 2025 年 2 月随论文发布于 arXiv,作者为北京大学与 DeepSeek)就是逐个击破这三道坎。而在进入 kernel 细节之前,先补一段背景:这个 kernel 是用 Triton 写的,为什么是 Triton?

铺垫:Triton 与 NSA 选它的理由#

Triton 是 OpenAI 开源的 GPU 编程领域特定语言(domain-specific language,DSL):用 Python 写 kernel 代码,由编译器翻译成针对目标 GPU 的 CUDA 代码。它对 CUDA 刚入门的读者来说,有几个需要建立的心智模型:

  • 程序(program)≈ CUDA 的线程块(thread block)。Triton 的 launch 配置里有一个 grid 参数(程序网格),每个 program 独立执行同一份 kernel 代码,靠 program id 区分自己处理哪一块数据。program 内部的数据并行由编译器自动映射到线程与向量化。
  • tile(块)是基本操作单位。程序员声明「我要加载一块 [M,N][M,N] 的矩阵、做一次矩阵乘、把结果写回」,对应 tl.loadtl.dottl.store 等原语。tl.dot 直接映射到 Tensor Core 的矩阵乘指令(如 mma/wgmma),编译器负责把 tile 切分到 warp 与寄存器。
  • 编译期常量(tl.constexprnum_warps / num_stages 是主要的调优旋钮:前者让块大小在编译期展开成常量(循环边界、索引计算全部特化),后者控制每个 program 的线程数与软件流水线深度。

手写 CUDA 实现稀疏注意力需要管理线程级的工作划分、共享内存(shared memory)布局、bank conflict 规避、异步拷贝等大量细节,而且换一代 GPU 架构往往要重写一遍。Triton 的 tile 抽象把「加载哪块、算哪块、写哪块」这类访存策略表达得比 CUDA 直观得多,编译器负责底层调度与架构适配。NSA 论文在 A100 上评测,对照组同样采用 Triton 实现的 FlashAttention-2——同后端对比,排除「Triton 比 CUDA 慢」这种不公平因素。

前向 kernel 总体设计:三个分支、三种处理方式#

论文把前向计算拆成两部分处理:

  • 压缩分支与滑窗分支直接复用 FlashAttention-2 风格的分块 kernel。论文原话是这两者的计算「与现成的 FlashAttention-2 kernel 天然兼容」。原因很直观:压缩分支的 key 序列是规则的——每 d=16d=16 个 token 生成一个新的压缩 key,压缩 key 排成一串连续的、随位置增长的序列,对压缩 key 的注意力就是一次普通的因果注意力(序列短得多);滑窗分支则是窗口内的普通因果注意力。两者都不涉及任何「动态选择」,FlashAttention 的连续分块 + 在线 softmax 流程原样适用。
  • 选择分支需要专门的 kernel,因为它的 KV 集合是按 query 动态变化的稀疏块索引集合。论文 3.4 节描述的正是这个 kernel,它的整体结构如论文 Figure 3 所示:

NSA 稀疏选择 kernel 的数据流组织(论文 Figure 3)
NSA 稀疏选择 kernel 的数据流组织(论文 Figure 3)

图片来源:NSA 论文 Figure 3。外层按 query 位置循环(Grid Loop,对应 Triton 的 program 网格),内层按该位置共享的稀疏块索引抓取 KV(Inner Loop);绿色块表示驻留在 SRAM 上的数据,蓝色块表示需要从 HBM 读取的数据。

这张图值得逐层读。最外层是 Grid Loop:kernel 的每个 program 负责序列上的一个 query 位置(准确说是「一个位置的整个 GQA 组」)。进入 program 后,先把该位置组内全部 query 头的数据从 HBM 读入 SRAM(图中绿色部分),再进入 Inner Loop:按照上一轮算好的共享块索引 It\mathcal{I}_{t},依次把被选中的 KV 块从 HBM(蓝色)读入 SRAM(绿色),每个块读入后立刻在 SRAM 上完成「该块对组内所有头的注意力累加」,然后进入下一个块。块索引向量 It\mathcal{I}_{t} 本身也是在每个 program 内预先算好的:压缩分支 kernel 算出分数后,经过组内求和与 top-n 排名得到。

论文将这套结构概括为三个特性,每个都对应着前面的一道坎。

特性一:Group-Centric Data Loading(组中心数据加载)。 对每个 query 位置 tt,把组内全部 hh 个头的 query 一次性载入 SRAM,记为 QR[h,dk]Q \in \mathbb{R}^{[h, d_{k}]},连同它们共享的稀疏块索引 It\mathcal{I}_{t}。这一步直接回应障碍二:既然组内所有头共享同一份 KV 块集合,就把「头」这个维度放进计算单元内部,让 KV 的每一次加载都能服务 hh 个头。与之相对,FlashAttention 的策略是「连续 query 块 × 全部 KV」,并行性来自 query 块维度;NSA 把并行性挪到了 GQA 组维度——两种策略是在「query 连续复用」与「KV 共享复用」之间做取舍,稀疏场景下后者才成立。

特性二:Shared KV Fetching(共享 KV 抓取)。 内层循环按 It\mathcal{I}_{t} 的顺序依次加载连续的 KV 块 KR[Bk,dk]K \in \mathbb{R}^{[B_{k}, d_{k}]}VR[Bk,dv]V \in \mathbb{R}^{[B_{k}, d_{v}]} 到 SRAM,其中 BkB_{k} 是 kernel 块大小,并且要求 BklB_{k} \mid l'BkB_{k} 整除选择块长 l=64l'=64)。两个细节都要展开。

「按索引顺序加载连续的 KV 块」对应障碍一:被选中的每个块在 HBM 里是连续的一段(每块 64 个 token 的 key 连续存放),kernel 按块为单位做合并访存(coalesced load),整块读入后驻留 SRAM 参与全部计算,绝不做 token 粒度的随机抓取——这正是 token 级稀疏方法(如第一篇批评过的 HashAttention)做不到的:它们每次读取跨越大段不相关数据,合并访存和 Tensor Core 的分块输入布局全部失效。

BklB_{k} \mid l'」则是一个对齐约束:选择块长度 64 被切成整数个 kernel 块,内层循环每次处理的大小固定,不会出现「一个选择块被切出一个不完整的尾巴」这种碎片化情况,循环边界可以完全特化。这个约束与压缩块的步长约束(第一篇讲过 dld \mid l)是同一个思路——让两个分块网格在硬件层面整齐对齐。

特性三:Outer Loop on Grid(外层循环交给 grid 调度)。 每个 query 位置的内层循环长度几乎完全相同——都约等于选中的块数 n=16n=16,与上下文长度无关。论文因此把 query 循环放在 Triton 的 grid 调度器上:一个 program 对应一个位置,所有 program 的工作量天然均衡。这回应了障碍三:不需要动态 worklist、原子计数器之类的负载均衡机制,静态调度即可让 SM 均匀吃饱。这也让 kernel 逻辑更简单——program 之间完全独立,没有跨 program 的同步或通信。

三个特性合起来,论文给出的评价是「实现了接近最优的算术强度」:组内共享消除了冗余 KV 传输(每个 KV 块只从 HBM 读一次),grid 外层循环让计算负载在 SM 间均衡分布。

算术强度账:组中心加载到底省了多少#

把上面的设计算一笔账,能直观看到「消除冗余 KV 传输」的分量。以论文效率评测的配置为例:每组 h=16h=16 个 query 头,dk=192d_{k}=192dv=128d_{v}=128。处理一个 Bk×dkB_{k}\times d_{k} 的 KV 块时:

  • 计算量(与头数相关):QK^T 约 2×h×Bk×dk2 \times h \times B_{k} \times d_{k} 次浮点运算,PV 约 2×h×Bk×dv2 \times h \times B_{k} \times d_{v} 次;
  • 若每个头各自加载 KV:HBM 流量是 h×Bk×(dk+dv)h \times B_{k} \times (d_{k}+d_{v})(同一条 KV 被 16 个头各读一遍);
  • 组中心加载后:HBM 流量是 Bk×(dk+dv)B_{k} \times (d_{k}+d_{v}),计算量不变。

也就是说,仅此一项设计就让该分支的算术强度提升了约 h=16h=16 倍——解码阶段是带宽受限的,这 16 倍直接对应解码带宽压力的 16 倍下降(选择分支内部)。这也从 kernel 的角度解释了第一篇强调的那句话:NSA 省的是内存访问,不是纸面上的 FLOPs。

解码与融合:三个分支在推理中如何协同#

训练和预填充阶段,三个分支各自是独立的注意力计算,最后按门控融合输出。解码阶段每个时间步只新增一个 query,此时每步需要加载的键值数量为(论文 5.2 节):

Nt=sld+nl+wN_{t}=\left\lfloor\frac{s-l}{d}\right\rfloor + nl' + w

其中 ss 是已缓存的序列长度,l=32l=32d=16d=16 是压缩块的块长与步长,n=16n=16l=64l'=64 是选择块的块数与块长,w=512w=512 是滑窗宽度。三项分别对应:压缩分支加载约 (sl)/d\lfloor(s-l)/d\rfloor 个压缩 key、选择分支加载 nl=1024nl'=1024 个原始 token、滑窗分支加载 512 个邻居 token。这三项里只有第一项随上下文线性增长,后两项与序列长度无关——这正是「上下文越长、稀疏率越高」的来源。

三个分支的输出用门控 gtcg_{t}^{c} 加权求和(第一篇已详细推导),门控参数在解码时每步照常计算。整个解码流程没有出现任何全注意力式的 KV 扫描:每步的 HBM 流量从「读一遍全部 ss 个 KV」变成「读 (sl)/d+nl+w\lfloor(s-l)/d\rfloor + nl' + w 个」。

反向 kernel:梯度怎么穿过稀疏结构#

论文报告反向传播在 64k 上下文下相对 Triton FlashAttention-2 加速 6.0 倍,但没有展开反向 kernel 的细节。结合第一篇关于可训练性的论证,反向的计算图是清晰的,这里把三点关键设计讲透。

第一,反向沿用与正向相同的稀疏数据流。 对选择分支,反向只需要访问被选中的 nlnl' 个 token 的 KV;对压缩分支,反向遍历压缩 key 序列(规则、连续)。与 FlashAttention 一样采用重计算(recomputation)策略:反向时重新计算注意力分数,而不是把前向的所有中间结果存下来。因为只选了 nn 个块,重计算的代价也只有全注意力的约 1/11.6,这正是反向能拿到 6× 加速的根本原因——反向的计算量与访存量同样随键数线性下降。

第二,梯度路径不经过 top-n。 第一篇论证过:top-n 只是内存寻址操作,不参与数值计算,因此反向时没有「对选择结果的梯度」。梯度沿压缩分支的 softmax 分数回传(压缩分支覆盖全序列、完全可微),对选中块 KV 的梯度只施加在被选中的块上,门控参数照常可导。训练时模型的「稀疏策略」通过压缩分支的分数梯度被端到端优化,选择分支只是忠实的执行者。

第三,工程实现上有经典的反向陷阱。 社区实现给出了两个可供参考的处理:tilde-research/nsa-release 支持选择注意力反向的 one-pass(单遍,原子操作)与 two-pass(两遍)两种变体;qflen/nsa-from-scratch 把 FlashAttention-2 式流式 softmax 反向改造成适配 gather 模式的版本,并把 dK/dVdK/dV 的梯度累加放在 fp32 缓冲里用原子加(atomic add)完成——多个 program 可能同时累加到同一块 KV 的梯度,普通写回会丢更新。这些细节在论文里读不到,但恰恰是「写一个能用的稀疏注意力反向」的真实工程面貌。

实测性能:前向 9×、反向 6×、解码 11.6×#

论文在 8 卡 A100 系统上评测(论文 5 节),模型配置沿用实验模型:GQA 组数 g=4g=4、每组 16 个头(h=16h=16)、dk=dq=192d_{k}=d_{q}=192dv=128d_{v}=128;NSA 参数 l=32l=32d=16d=16l=64l'=64n=16n=16w=512w=512。对照实现是 Triton 版 FlashAttention-2,保证「同后端公平对比」。结果如论文 Figure 6:

NSA kernel 与 Triton FlashAttention-2 的延迟对比(论文 Figure 6)
NSA kernel 与 Triton FlashAttention-2 的延迟对比(论文 Figure 6)

图片来源:NSA 论文 Figure 6。左图:不同上下文长度下 Triton FlashAttention-2(全注意力)与 NSA 的 kernel 延迟(毫秒),差距随序列变长持续拉大;右图:64k 下的加速比,前向 9.0 倍、反向 6.0 倍。

左图的横轴是上下文长度(8k 到 64k),两条曲线分别对应全注意力 kernel 与 NSA kernel 的端到端延迟:8k 时差距已经可观,到 64k 时差距扩大到近一个数量级。论文给出的峰值数据:64k 上下文下前向 9.0×、反向 6.0×。加速比随上下文变长的单调上升,直接对应键数公式里「只有压缩项线性增长」的结构。

解码侧(论文 Table 4)的账本如下——注意这里的「键数」同时就是「每步 HBM 加载的 token 数」,因为解码是带宽受限的:

上下文长度全注意力键数NSA 键数期望加速比
8k819220484.0×
16k1638425606.4×
32k3276835849.1×
64k65536563211.6×

论文在 5.2 节明确说明:解码阶段算术强度极低、带宽受限,期望加速比与访存量近似线性——所以 64k 下访存降到 1/11.6,解码延迟理论上也接近 1/11.6。这个数字与 Table 4 的期望加速比完全一致。

论文还解释了加速的两个来源(5.1 节):一是块状内存访问通过合并加载最大化了 Tensor Core 利用率;二是 kernel 内精细的循环调度消除了冗余 KV 传输——正是上一节三个特性的直接收益。

一个值得对照的社区数据点:qflen/nsa-from-scratch(从零重写 NSA 的社区项目,选择分支用 Triton + Hopper WGMMA 指令实现)在其基准中报告,64k 上下文下 NSA 前向吞吐比 FlashAttention-3 快 7.07 倍(10.84M vs 1.53M token/s,H100 NVL)。这是社区实现的自报数据,测试协议、数值精度设定与论文不完全一致,但方向上一致地说明:当 kernel 把访存组织对之后,稀疏注意力的加速是实打实能拿到的。该项目的 README 还记录了几个实操层面的发现:top-k 选块的并列次序(tie-breaking)会影响数值稳定性与性能、gather 到 tile 后的因果掩码需要小心处理、以及 48 组配置的自动调优(autotune)对 BLOCK_M/BLOCK_N/num_warps/num_stages 的敏感性——这些是论文没有展开、但任何复现者都会踩到的工程细节。

27B 模型:从预训练到长上下文的完整实验#

训练设置#

实验模型是一个 27B 总参数、3B 激活参数的 MoE(Mixture-of-Experts,混合专家)模型:30 层、隐藏维 2560;注意力采用 GQA,4 组共 64 个头,dq=dk=192d_{q}=d_{k}=192dv=128d_{v}=128;MoE 采用 DeepSeekMoE 结构,72 个路由专家加 2 个共享专家,top-k 取 6;为保证训练稳定,第一层的 MoE 被替换成 SwiGLU 形式的 MLP(论文 4.1 节)。

训练流程:全注意力基线与 NSA 都先在 8k 长度的文本上预训练 270B tokens,再用 YaRN 把上下文扩展到 32k,做续训与监督微调(supervised fine-tuning,SFT),两边都训练到完全收敛以保证公平。训练损失曲线如论文 Figure 4:

27B 模型预训练损失对比(论文 Figure 4)
27B 模型预训练损失对比(论文 Figure 4)

图片来源:NSA 论文 Figure 4。全注意力基线与 NSA 的预训练损失都平稳下降,且 NSA 全程略低于全注意力基线。

稀疏模型在预训练中损失不升反降,本身就是一个强信号:稀疏性没有成为模型的负担,反而起到了某种正则化/聚焦作用(论文在通用基准一节给出了解释:强制模型聚焦最重要信息,相当于过滤噪声注意力路径)。

通用基准:9 项里赢 7 项#

预训练完成后在知识、推理、代码三类共 9 项基准上评测(论文 Table 1):

基准MMLUMMLU-PROCMMLUBBHGSM8KMATHDROPMBPPHumanEval平均
Full Attention0.5670.2790.5760.4970.4860.2630.5030.4820.3350.443
NSA0.5650.2860.5870.5210.5200.2640.5450.4660.3480.456

NSA 在 9 项中赢 7 项,平均分 0.456 对 0.443。输的两项(MMLU 差 0.002、MBPP 差 0.016)都在噪声范围内,而赢的项目里推理类增益最突出:DROP +0.042、GSM8K +0.034。论文对这个现象的解释值得留意:稀疏预训练让模型被迫把注意力集中在最重要的信息上,等于把无关注意力路径的噪声过滤掉了——这正是一篇论文主张「原生稀疏」的核心理据之一:稀疏不只是省钱的代价,它本身可能是更好的归纳偏置。

长上下文:64k 检索满分与 LongBench 第一#

长上下文能力先看检索。论文 Figure 5 展示了 64k 上下文的 needle-in-a-haystack(大海捞针)测试:

64k 上下文大海捞针测试(论文 Figure 5)
64k 上下文大海捞针测试(论文 Figure 5)

图片来源:NSA 论文 Figure 5。横轴是针所在的位置,NSA 在 64k 上下文的所有位置都达到 100% 检索准确率(图中全绿)。

满分的原因正是分层设计本身:压缩分支以极低代价完成全局扫描,把「针大概在哪一段」定位出来;选择分支随即把该段读细,找到精确位置。粗粒度定位 + 细粒度确认的分工,天然适合检索类任务。

LongBench 的对比更有信息量(论文 Table 2)。为保证稀疏度对齐,所有稀疏方法(H2O、infLLM、Quest、Exact-Top)都被限制为每个 query 激活 2560 个 token——对应 NSA 在 32k 序列时的平均激活数;按 StreamingLLM 的做法,这 2560 个 token 的预算里包含序列开头的 128 个 token 和最近的 512 个局部 token。在这个对齐预算下,NSA 的平均分 0.469 高于所有基线:比全注意力(0.437)高 0.032,比 Exact-Top(0.423,理论上的「完美的 token 级 oracle 选择」)高 0.046。

值得注意的几点:Exact-Top 是「先算全注意力、再按真实分数挑 top-n」的理想化方法——它拿到了 oracle 级别的选择质量,但 LongBench 上反而低于 NSA。这说明稀疏方法的质量上限不只在「选得准不准」,还在于选择与模型是否协同演化:NSA 的选择结构在预训练中与模型一起被优化,模型学会了适应「只被选中块看到」的注意力模式;Exact-Top 是事后施加的选择,模型并没有为此做过任何适配。在具体子集上,NSA 在多跳问答(HPQ +0.087、2Wiki +0.051)、代码理解(LCC +0.069)、段落检索(PassR-en +0.075)上的增益尤其大。论文也如实报告了方法学上的取舍:部分子集因为所有模型的分数都过低、缺乏区分度而被排除。

思维链推理:AIME 上的显著领先#

为了验证 NSA 与进阶训练范式的兼容性,论文做了 DeepSeek-R1 蒸馏实验:用 R1 生成的 10B tokens、32k 长度的数学推理轨迹对两个模型做监督微调,得到 Full Attention-R 与 NSA-R,然后在 AIME 24 上评测(温度 0.7、top-p 0.95、每题采样 16 次取平均)。结果(论文 Table 3):

生成长度上限819216384
Full Attention-R0.0460.092
NSA-R0.1210.146

8k 上限下 NSA-R 是基线的 2.6 倍(+0.075),16k 下优势保持(+0.054)。论文给出的解释是:原生训练的稀疏模式善于抓住长距离逻辑依赖——数学推理恰恰需要跨很长的 token 距离追踪推导链条;同时稀疏架构在推理深度增长时维持了足够的上下文密度,没有出现能力崩坏。这个实验的意义在于:稀疏注意力不是「省钱但伤能力」的妥协,在推理模型(thinking model)场景下反而可能成为优势。

消融实验:为什么辅助损失与启发式选择都不行#

论文 6.1 节和 Figure 7 用 3B 模型做了三种选择策略的消融对比(训练损失曲线):

3B 模型不同选择策略的消融(论文 Figure 7)
3B 模型不同选择策略的消融(论文 Figure 7)

图片来源:NSA 论文 Figure 7。Full Attention 与 NSA 的损失接近且最低;辅助损失式选择(近似 SeerAttention 思路)与启发式无参数选择(Quest 式)的损失都明显偏高。

  • 辅助损失式选择(近似 SeerAttention 的思路):额外引入查询与块代表 key 来预测块重要性,块级监督信号来自「块内注意力分数按均值池化」,用 KL 散度约束预测。论文发现这类方法有两个问题:额外的打分算子增加开销;更重要的是辅助目标与主目标(语言建模损失)不完全一致,优化的轨迹被带歪。
  • 启发式无参数选择(Quest 式):用 query 与 key 分块的逐维 min-max 做点积打分。召回率低——min-max 估计的块分数与真实注意力分布偏差大,选出来的块常常不是真正重要的块。论文还尝试了冷启动变体(前 1000 步先跑全注意力再切换),同样不行。
  • NSA:损失与全注意力基线持平甚至略低。

图上的对比直接支撑了第一篇的核心论断:打分器必须与主模型端到端联合训练,任何「事后可微」或「启发式」的替代方案都在损失曲线上付出代价。论文 6.1 节还讨论了聚类式选择(ClusterKV 式)在训练场景的三个工程障碍:动态聚类本身的计算开销;簇间负载不均衡(在 MoE 系统的专家并行(expert parallelism,EP)下,各专家组执行时间本就参差,叠加聚类不均会形成持续负载失衡);以及必须周期性重聚类、且要求分块顺序训练协议的实现约束。这些障碍都属于「理论上可行、工程上难以为继」。

另一个支撑块级稀疏合理性的观察来自注意力分布本身。论文 Figure 8 画了 27B 全注意力模型各层的注意力图:

全注意力模型的注意力分布可视化(论文 Figure 8)
全注意力模型的注意力分布可视化(论文 Figure 8)

图片来源:NSA 论文 Figure 8。亮色区域表示高注意力值,可见高分区域呈明显的块状聚集(blockwise clustering):相邻 key 的注意力分数往往相近。

这张图回答了一个基础问题:为什么按「块」选而不是按 token 选是安全的。如果高分位置在序列上随机散布,块级选择会漏掉大量信息;实测的块状聚集说明「重要的 key 往往聚在一起」,块级稀疏丢的信息有限,却换来了硬件友好的连续访存——与 kernel 设计形成了「数据分布证据 → 算法结构 → 硬件实现」的完整闭环。

从 NSA 到 DSA:DeepSeek-V3.2 的工程化落地#

2025 年 9 月 29 日,DeepSeek 开源了 DeepSeek-V3.2-Exp——基于 V3.1-Terminus 的 128K 长上下文模型,其唯一的架构改动就是引入了 DSA(DeepSeek Sparse Attention,DeepSeek 稀疏注意力),并随开源发布了一份 6 页的技术报告(DeepSeek_V3_2.pdf)。这是 NSA 思想第一次走进旗舰级生产模型,但 DSA 不是 NSA 的照搬,它在三个关键点上做了演进。

架构:闪电索引器 + 细粒度 token 选择#

DSA 的口号是「先过滤、后计算」(filter first, compute later)。它的两个组件:一个轻量级的闪电索引器(lightning indexer)负责给每个 query 与历史 token 的相关性打分,一个细粒度 token 选择机制负责挑出分数最高的 token 参与正式注意力。索引分数定义为:

It,s=j=1HIwt,jIReLU(qt,jIksI)I_{t,s}=\sum_{j=1}^{H_{I}}w_{t,j}^{I}\cdot\operatorname{ReLU}\left(q_{t,j}^{I}\cdot k_{s}^{I}\right)

HIH_{I} 是索引器头数,qt,jIRdIq_{t,j}^{I}\in\mathbb{R}^{d_{I}} 与标量权重 wt,jIw_{t,j}^{I} 由 query token 的隐藏向量推导,ksIk_{s}^{I} 由历史 token 推导。逐项解释:每个索引头对 query 与 key 做一次内积,ReLU 截掉负分,再用每头一个可学习权重加权求和。选 ReLU 而不是 softmax 是明确的吞吐考量——softmax 需要跨序列维的归约与指数运算,ReLU 是逐元素的;加上索引器头数少、可以用 FP8 低精度实现,整个打分过程的计算开销被压得很低。选出的 top-k 键值条目参与正式注意力:

ut=Attn(ht,  {csIt,sTop-k(It,:)})u_{t}=\operatorname{Attn}\left(h_{t},\;\left\{c_{s}\mid I_{t,s}\in\operatorname{Top\text{-}k}(I_{t,:})\right\}\right)

与 NSA 的关键区别在粒度:NSA 按选(64-token 块 × 16 块 = 1024 个 token),DSA 按token选,每个 query 激活 k=2048 个键值 token(Hugging Face 配置 index_topk=2048;索引器配置为 64 个头、头维 128,前 3 层保持稠密注意力 first_k_dense_replace=3)。报告称这是「首次实现细粒度稀疏注意力」。另一个工程决策是把它实例化在 MLA(Multi-head Latent Attention,多头潜在注意力,见站内MLA 完全拆解)的 MQA 模式下:每个潜在向量(MLA 的键值条目)被该 query 的所有注意力头共享——报告明确引用了 NSA 论文的结论:kernel 层面,每个键值条目必须被多个 query 共享,计算效率才成立。这与 NSA kernel 的组中心加载是同一个原则在 MLA 语境下的延伸。整体架构如图 1 所示:

DeepSeek-V3.2-Exp 的 DSA 架构(技术报告 Figure 1)
DeepSeek-V3.2-Exp 的 DSA 架构(技术报告 Figure 1)

图片来源:DeepSeek-V3.2-Exp 技术报告 Figure 1。绿色部分展示了 DSA 的工作方式:隐藏向量 hth_{t} 分出索引器分支(lightning indexer),对历史 KV 打分后由 top-k 选择器选出键值条目,核心注意力(core attention,MLA 的 MQA 模式)只对选中的条目计算。

训练:两阶段的继续训练#

DSA 是「改造」而非「从零训练」,这决定了它与 NSA 在方法论上的分工。报告的训练方案分两个阶段(报告 2.1 节):

  • 稠密热身阶段(dense warm-up):保持稠密注意力,冻结除索引器外的全部参数,只训练索引器。训练目标是把索引器的分布对齐到真实注意力的分布:把主注意力分数按头求和、沿序列维做 L1 归一化得到目标分布 pt,:p_{t,:},用 KL 散度做损失。学习率 10310^{-3},只训 1000 步,每步 16 条 128K 序列,合计 2.1B tokens——热身成本极低,只为了让索引器「学会像全注意力一样思考」。
  • 稀疏训练阶段:放开全部参数,引入 top-k 选择,让模型整体适应稀疏模式。索引器从计算图 detach(脱离),只由 KL 损失优化、且只在被选中的 token 集合上计算;主模型只由语言建模损失优化。学习率 7.3×1067.3\times10^{-6},k=2048,15000 步、每步 480 条 128K 序列,合计 943.7B tokens。之后的后训练(专家蒸馏 + 混合 GRPO 强化学习)沿用 V3.1-Terminus 的完整流程,保证对比口径一致。

对照 NSA:NSA 是「从零预训练时稀疏原生生效」,DSA 是「对已收敛的稠密模型做稀疏化续训」。两条路线各有适用场景,DSA 的路线让 DeepSeek 不必重训 671B 级模型就能吃到稀疏注意力的收益——这正是它把 NSA 的「原生可训练」从论文理念推进到生产约束(复用已有检查点)的方式。

效果与成本#

基准层面,报告与 README 给出的结论是「与 V3.1-Terminus 持平」:MMLU-Pro 85.0/85.0、GPQA-Diamond 80.7/79.9、SWE Verified 68.4/67.8 大体相当,AIME 2025 从 88.4 升到 89.3、Codeforces 从 2046 升到 2121,Humanity’s Last Exam 从 21.7 降到 19.8——有升有降,总体在噪声带内。成本层面,技术报告的 Figure 3 展示了 H800 集群(按 2 美元/GPU·小时租赁价折算)上每百万 token 的推理成本随上下文长度的变化:

DSA 的推理成本曲线(技术报告 Figure 3)
DSA 的推理成本曲线(技术报告 Figure 3)

图片来源:DeepSeek-V3.2-Exp 技术报告 Figure 3。左图预填充(prefilling)、右图解码(decoding):每百万 token 成本随上下文长度变化,V3.2-Exp 的曲线在长上下文区间明显低于 V3.1-Terminus。

两条曲线的分叉趋势一致:上下文越长,稀疏注意力的成本优势越大——与 NSA 论文「加速比随长度单调上升」的结论一脉相承。报告还提到一个短上下文侧的工程细节:短序列预填充时专门实现了 masked MHA 模式来模拟 DSA,在短上下文下获得更高效率。发布当日 DeepSeek 即下调 API 定价,多家报道称部分档位降幅超过 50%。

开源层面,DSA 的 kernel 拆分值得注意:索引器 logit kernel(含分页版本)放进 DeepGEMM,稀疏注意力 kernel 放进 FlashMLA——复用两个成熟内核库,而不是另起炉灶。此外,2025 年 11 月 17 日官方发布了一项修复:推理示例代码里索引器模块的 RoPE 布局与 MLA 模块不一致——索引器输入要求非交错(non-interleaved)布局,而 MLA 的 RoPE 期望交错(interleaved)布局,该差异可能损害模型性能。这个「算法正确、实现错位」的修复公告本身就是生产级工程的一部分:稀疏注意力在真实服务栈里的每一层(数据布局、kernel、调度)都可能藏着 bug。

NSA 与 DSA 的对照#

维度NSA(论文,2025.02)DSA(V3.2-Exp,2025.09)
稀疏粒度块级:64-token 块 × 16 块 = 1024 token/querytoken 级:top-2048 个 token/query
选择依据压缩分支 softmax 分数(副产品,零额外开销)独立的闪电索引器(FP8、ReLU,额外训练)
打分器训练随主模型端到端梯度训练稠密热身(KL 蒸馏)+ 稀疏阶段 KL 对齐,detach 优化
训练方式从零预训练(270B tokens)在 V3.1-Terminus 上继续训练(约 946B tokens)
主干注意力标准 MHA/GQAMLA 的 MQA 模式
kernel 形态Triton(论文)+ 社区实现CUDA:DeepGEMM(索引器)+ FlashMLA(稀疏注意力)
定位方法论:证明原生稀疏可训练且不损质量工程化:旗舰模型上的生产验证

这张表也揭示了稀疏注意力研究的主线:NSA 证明「稀疏可以原生训练、可以拿到 9-11 倍的实测加速」,DSA 则回答「如何在不动 671B 级检查点的前提下把这套机制搬进生产模型」。两条路线在「打分器与主模型的协同训练」这一核心原则上完全一致。

小结#

本篇把 NSA 从算法推进到了工程与实验层,核心结论可以收拢为四点。

第一,kernel 是理论加速的兑现层。NSA 用「组中心数据加载 + 共享 KV 抓取 + grid 外层循环」三个设计,分别击破稀疏访存、GQA 共享带宽和 SM 负载均衡三道坎:KV 块只从 HBM 读一次、合并连续访问、program 间静态均衡。第二,实测数据验证了理论:64k 上下文、8×A100 上,前向 9.0×、反向 6.0×、解码 11.6×,加速比随上下文变长单调上升;社区从零实现(如 nsa-from-scratch)在 H100 上对 FA-3 的 7.07× 前向加速从另一个硬件世代印证了这一点。第三,27B 模型实验证明稀疏不伤质量:9 项通用基准赢 7 项、LongBench 平均分领先全注意力与 oracle 式 Exact-Top、64k 大海捞针满分、R1 蒸馏后 AIME 成绩翻倍——消融实验进一步说明辅助损失与启发式打分都打不过与主模型联合训练的选择器。第四,DSA 是这套思想的工程化终点:token 级细粒度选择 + 独立闪电索引器 + 续训改造,让 671B 级生产模型在长上下文下把成本曲线显著压低。

至此,NSA 系列的两篇覆盖了它的全部层次:算法与可训练性原理(第一篇)、kernel 实现与实验验证、以及生产落地(本篇)。稀疏注意力的其他路线——如 MoBA 的 block-wise 预计算路由、MInference 的静态稀疏预填充等——与 NSA 的对比与融合,可以作为后续单独的话题展开。

参考资料#

  1. Native Sparse Attention: Hardware-Aligned and Natively Trainable Sparse Attention(NSA 论文 arXiv 页)
  2. NSA 论文 arXiv HTML 版(本文 Figure 3/5 配图来源)
  3. ACL Anthology 正式版本(ACL 2025 Best Paper)
  4. DeepSeek-V3.2-Exp 开源仓库(DSA 技术报告、基准表与 kernel 链接)
  5. DeepSeek-V3.2-Exp 技术报告 PDF(DeepSeek Sparse Attention 架构与训练细节)
  6. DeepSeek Open-Sources V3.2-Exp and Unveils New Sparse Attention Mechanism DSA(36Kr 报道)
  7. Hugging Face Transformers DeepSeek-V3.2 模型文档(index_topk 等 DSA 配置)
  8. tilde-research/nsa-release:NSA 的 PyTorch+Triton+FlexAttention 实现
  9. qflen/nsa-from-scratch:从零重写 NSA(Triton + Hopper WGMMA,含性能与数值细节的完整记录)
  10. lucidrains/native-sparse-attention-pytorch:NSA 的可安装研究实现
  11. fla-org/flash-linear-attention:线性注意力库(含并行 NSA 实现)
  12. DeepGEMM Pull Request #200(DSA 索引器 logit kernel,含分页版本)
  13. FlashMLA Pull Request #98(DSA 稀疏注意力 kernel)
  14. NSA 详解:Compression + Selection + Sliding Window(yudonglee.me 技术博客)

文章分享

如果这篇文章对你有帮助,欢迎分享给更多人!

NSA 原生稀疏注意力(二):Triton Kernel 设计与 27B 模型实验
https://pinghaoyang.com.cn/aigc/posts/native-sparse-attention-part-2/
作者
平昊阳
发布于
2026-08-28
许可协议
CC BY-NC-SA 4.0

评论区

Profile Image of the Author
平昊阳
乘长风,破巨浪, 展鸿图于未央!
--
总访问量
--
访客数
公告
欢迎来到我的个人博客!欢迎关注交流吖!
更多相关公告,见
社交-留言」。
音乐
封面

音乐

暂未播放

0:000:00
暂无歌词
站点统计
文章
80
分类
18
标签
105
总字数
654,628
运行时长
0
最后活动
0 天前

文章目录