DeepSeek 提出新一代注意力机制 Native Sparse Attention (NSA): 硬件对齐且本地可训练的稀疏注意力机制
arXiv | https://arxiv.org/abs/2502.11089
摘要:
长上下文建模对于下一代语言模型至关重要,然而标准注意力机制的高计算成本带来了重大的计算挑战。 稀疏注意力为提高效率同时保持模型能力提供了有前景的方向。我们提出了 NSA,一种本地可训练的稀疏注意力机制,将算法创新与硬件对齐的优化相结合,以实现高效的长上下文建模。NSA采用动态分层稀疏策略,结合粗粒度的 token 压缩与细粒度的 token 选择,以保留全局上下文意识和局部精度。
在稀疏注意力设计方面实现了两项关键创新:(1)我们通过算术强度平衡的算法设计实现显著的加速,并针对现代硬件进行实现优化。(2)我们实现了端到端的训练,减少了预训练计算量而不牺牲模型性能。

一、引言
学术研究社区越来越多地认识到,长上下文建模是下一代大规模语言模型的关键能力,这推动了从深度推理、仓库级代码生成到多轮自主代理系统的多样化实际应用。近期的突破包括 OpenAI o 系列模型、DeepSeek-R1 和 Gemini 1.5 Pro ,这些模型能够处理整个代码库、长文档,并在数千个 token 的范围内保持连贯的多轮对话,同时进行长距离依赖的复杂推理。然而,基础的注意力机制随着序列长度的增加而表现出的高复杂性成为了一个关键的延迟瓶颈,理论估算表明,在解码 64k 长度的上下文时,基于 softmax 架构的注意力计算占据了总延迟的 70-80%,这凸显了迫切需要更高效的注意力机制。
一种高效处理长上下文建模的自然方法是利用 softmax 注意力的固有稀疏性,通过选择性地计算关键的查询-键对,可以显著减少计算开销同时保持性能。但现有的稀疏注意力方法在实际部署中往往表现不佳,未能实现与理论预期相当的速度提升;此外,大多数方法主要关注推理阶段,缺乏有效的训练支持,无法充分利用注意力的稀疏模式。
为了解决这些限制,有效的稀疏注意力机制的部署必须应对两大关键挑战:
- **硬件对齐的推理加速:**将理论上的计算减少转化为实际的加速需要在预填充和解码阶段进行硬件友好的算法设计,以缓解内存访问和硬件调度瓶颈;
- **训练感知的算法设计:**通过使用可训练的操作符实现端到端计算,从而降低训练成本同时保持模型性能。这些要求对于实现快速长上下文推理或训练至关重要。
为了实现更为有效和高效的稀疏注意力机制,我们提出了本地稀疏注意力机制 NSA(Natively Trainable Sparse Attention),该架构整合了层级化 token 建模。

左图:通过三个并行的注意力分支处理输入序列,对于给定的查询,
- 前置的键和值:粗粒度的 token 压缩
- 重要的 token 块:细粒度的 token 选择
- 局部上下文:滑动注意力窗口
**右图:**每个分支生成的不同注意力模式的可视化。
- 绿色区域表示需要计算注意力分数的区域
- 白色区域表示可以跳过的区域
NSA 引入了两个核心创新,以满足上述关键要求:
-
**硬件对齐系统:**优化块状稀疏注意力机制以提高张量核利用率和内存访问效率,确保算术强度平衡。
-
**训练感知设计:**通过高效算法和反向操作实现端到端训练的稳定性。这种优化使NSA能够支持高效的部署和端到端训练。
二、重新思考稀疏注意力机制
现代稀疏注意力方法在降低 transformer 模型的理论计算复杂度方面取得了显著进展。然而,大多数方法主要在推理过程中应用稀疏性,而保留了一个预训练的全注意力主干,这可能会引入架构偏见,从而限制了它们充分利用稀疏注意力优势的能力,本节通过两个关键视角系统地分析了这些局限性。
2.1 有效推理的错觉
尽管在注意力计算中实现了稀疏性,许多方法在推理延迟上未能实现相应的减少,主要原因在于两个挑战:
2.1.1 相位受限的稀疏性
- H2O 等方法在自回归解码过程中应用稀疏性,但在填充阶段需要进行计算密集型的预处理(如注意力图计算、索引构建)。
- MInference 等方法仅关注填充阶段的稀疏性。这些方法在所有推理阶段均未能实现加速,因为至少一个阶段的计算成本仍与全注意力相当。
相位专业化降低了这些方法在以填充为主的工作负载(如书籍摘要和代码补全)或以解码为主的工作负载(如长链条推理)中的加速能力。
2.1.2 与高级注意力架构不兼容
某些稀疏注意力机制无法适应现代解码高效架构,如多查询注意力(MQA)和组查询注意力(GQA)。MQA 和 GQA 通过在多个查询头之间共享 KV 显著减少了解码过程中的内存访问瓶颈。
- 在类似于 Quest 的方法中,每个注意力头独立选择其自己的 KV 缓存子集。尽管这种方法在多头注意力(MHA)模型中展示了持续的计算稀疏性和内存访问稀疏性,但在基于 GQA 等架构的模型中 KV 缓存的内存访问量对应于同一 GQA 组内所有查询头选择的并集。
尽管某些稀疏注意力方法可以减少计算量,但分散的内存访问模式与高级架构的高效内存访问设计相冲突。
2.2 可训练稀疏性的神话
我们对本地可训练的稀疏注意力机制的追求源于分析仅推理方法时的两个关键洞察:
- **性能退化:**事后应用稀疏性迫使模型偏离其预训练优化轨迹。Chen et al. (2024) 展示了前20%的注意力机制仅能覆盖总注意力分数的70%,这使得预训练模型中的检索头在推理过程中容易受到剪枝的影响。
- **训练效率需求:**高效处理长序列训练对于现代大语言模型的发展至关重要,包括在更长文档上的预训练以增强模型容量,以及随后的适应阶段,如长上下文微调和强化学习。
然而,现有的稀疏注意力方法主要针对推理,而对训练过程中的计算挑战则基本未予解决。这一限制阻碍了通过高效训练开发更强大长上下文模型的进程,但将现有稀疏注意力机制适应于训练存在诸多挑战:
- **非训练组件。**在如 ClusterKV(包括 k-means 聚类)和 MagicPIG(包括基于 SimHash 的选择)等方法中,离散操作会在计算图中产生不连续性。这些不可训练组件阻碍了梯度在 token 选择过程中的流动,限制了模型学习最优稀疏模式的能力。
- 低效的反向传播。一些理论上可训练的稀疏注意力方法在实践中面临着训练效率低下的问题。如 HashAttention 等方法采用的以 token 为单位的选择策略,在注意力计算过程中需要从 KV 缓存中加载大量独立的词元,导致了非连续的内存访问。这种非连续的内存访问方式妨碍了 FlashAttention 等依赖连续内存访问和块状计算以实现高吞吐量的快速注意力技术的有效适应。因此,实现方式被迫采用低硬件利用率的方法,显著降低了训练效率。
三、方法
3.1 背景
3.1.1 注意力机制
注意力机制在语言建模中广泛应用,其中每个查询 token qt\mathbf{q}_tqt 会计算与所有先前的键 token k:t\mathbf{k}_{:t}k:t 的相关性得分以生成值 token v:t\mathbf{v}_{:t}v:t 的加权和。对于长度为 ttt 的输入序列,注意力操作定义为:
ot=Attn(qt,k:t,v:t)=∑i=1tαt,ivi∑j=1tαt,j,αt,i=eqtTkidk
\mathbf{o}_t=\text{Attn}(\mathbf{q}_t, \mathbf{k}_{:t}, \mathbf{v}_{:t})=\sum_{i=1}^t\frac{\alpha_{t,i}\mathbf{v}_{i}}{\sum_{j=1}^t\alpha_{t,j}},\alpha_{t,i}=e^{\frac{\mathbf{q}_t^\mathrm T \mathbf{k}_{i}}{\sqrt{d_k}}}
ot=Attn(qt,k:t,v:t)=i=1∑t∑j=1tαt,jαt,ivi,αt,i=edkqtTki
αt,i\alpha_{t,i}αt,i 表示 qt\mathbf{q}_tqt 和 ki\mathbf{k}_{i}ki 之间的注意力权重,dkd_kdk 是键的特征维度。随着序列长度的增加,注意力计算在总体计算成本中所占的比例越来越大,这为长上下文处理带来了重大挑战。
3.1.2 算术强度
算术强度是指计算操作与内存访问的比例,内在地决定了在硬件上的算法优化。每个GPU的关键算术强度由其峰值计算能力与内存带宽决定,计算方法是将这两项硬件限制进行比值运算。对于计算任务而言,算术强度高于这一关键阈值时会成为计算绑定型(受限于GPU FLOPS),而低于这一阈值时则会成为内存绑定型(受限于内存带宽)。
-
在因果自注意力机制中,训练和预填充阶段的批量矩阵乘法和注意力计算表现出高的算术强度,使得这些阶段在现代加速器上主要受计算约束。
-
自回归解码则因每次前向传递生成一个token,但需要加载整个键值缓存,导致算术强度较低,从而变得主要受内存带宽约束。 这导致了不同的优化目标:在训练和预填充阶段减少计算成本,而在解码阶段减少内存访问。
3.2 整体框架
为了充分发挥自然稀疏模式下注意力机制的潜力,我们提出用每个查询 qt\mathbf{q}_tqt 对应的一组更紧凑且信息密集的表示键值对 K~t,V~t\tilde{K}_t,\tilde{V}_tK~t,V~t 替换原始的键值对 k:t,v:t\mathbf{k}_{:t}, \mathbf{v}_{:t}k:t,v:t:
K~t=fK(qt,k:t,v:t),V~t=fV(qt,k:t,v:t)ot∗=Attn(qt,K~t,V~t)
\tilde{K}_t=f_K(\mathbf{q}_t, \mathbf{k}_{:t}, \mathbf{v}_{:t}),\tilde{V}_t=f_V(\mathbf{q}_t, \mathbf{k}_{:t}, \mathbf{v}_{:t})\\
\mathbf{o}^*_t=\text{Attn}(\mathbf{q}_t, \tilde{K}_t,\tilde{V}_t)
K~t=fK(qt,k:t,v:t),V~t=fV(qt,k:t,v:t)ot∗=Attn(qt,K~t,V~t)
其中,K~t,V~t\tilde{K}_t,\tilde{V}_tK~t,V~t 是基于当前查询 qt\mathbf{q}_tqt 和上下文记忆 k:t,v:t\mathbf{k}_{:t}, \mathbf{v}_{:t}k:t,v:t 动态构建的,可以设计各种映射策略以获得不同类别的 K~tc,V~tc\tilde{K}_t^c,\tilde{V}_t^cK~tc,V~tc,并将其组合如下:
ot∗=∑c∈Cgtc⋅Attn(qt,K~tc,V~tc)
\mathbf{o}^*_t=\sum_{c\in C}g_t^c\cdot\text{Attn}(\mathbf{q}_t, \tilde{K}_t^c,\tilde{V}_t^c)
ot∗=c∈C∑gtc⋅Attn(qt,K~tc,V~tc)
NSA 有三种映射策略 C={cmp,slc,win}C = \{\text{cmp}, \text{slc}, \text{win}\}C={cmp,slc,win},分别代表 token 压缩、token 选择和滑动窗口策略,用于键和值。对于相应的策略 ccc,门控分数 gtc∈[0,1]g_t^c\in[0,1]gtc∈[0,1] 是通过一个 MLP 和 sigmoid 激活从输入特征中导出的。
令 NtN_tNt 表示重新映射的键和值的总数:
Nt=∑c∈Csize[K~tc]
N_t=\sum_{c\in C}\text{size}[\tilde{K}_t^c]
Nt=c∈C∑size[K~tc]
通过确保 Nt≪tN_t\ll tNt≪t 来保持高稀疏比率。
3.3 算法设计
3.3.1 Token 压缩
通过将连续的键或值块聚合为块级表示获得压缩的键和值,这些压缩的键和值捕捉整个块的信息。压缩键表示为:
K~tcmp=fKcmp(k:t)={φ(kid+1:id+l)∣1≤i≤⌊t−ld⌋}
\tilde{K}_t^{\text{cmp}}=f_K^{\text{cmp}}(\mathbf{k}_{:t})=\{\varphi(\mathbf{k}_{id+1:id+l})|1\le i\le \lfloor \frac{t-l}{d} \rfloor\}
K~tcmp=fKcmp(k:t)={φ(kid+1:id+l)∣1≤i≤⌊dt−l⌋}
其中 lll 表示块的长度,ddd 表示相邻块之间的滑动步长,而 φ\varphiφ 是一个带有内部块位置编码的可学习 MLP,用于将块中的键映射到单个压缩键,K~tcmp∈Rdk×⌊t−ld⌋\tilde{K}^{\text{cmp}}_t \in \mathbb{R}^{d_k \times \lfloor \frac{t-l}{d} \rfloor}K~tcmp∈Rdk×⌊dt−l⌋ 是由压缩键组成的张量,通常采用 d<ld\lt ld<l 来减轻信息碎片化的问题。压缩值同理。
压缩表示捕获了更粗粒度的高层语义信息,并减少了注意力机制的计算负担。
3.3.2 Token 选择
仅使用压缩键时,值可能会丢失重要的细粒度信息,这促使我们选择性地保留个别键和值。因此我们还使用了一种高效的 token 选择机制,能够以较低的计算开销识别并保留最相关的令牌。
区块选择
选择策略以连续的空间区块处理键和值序列,这主要是出于两个关键因素的考虑:
- 硬件效率:区块选择对于在现代 GPU 上实现高效计算至关重要,因为现代 GPU 架构在连续区块访问方面表现出比随机索引读取更高的吞吐量。此外,区块计算能够最大化张量内核的利用效率。这种架构特性已经确立了区块化内存访问和计算作为高性能注意力实现的基本原则。
- 注意力分数的固有分布模式:区块选择遵循注意力分数的固有分布模式。注意力分数往往表现出空间连续性,这意味着相邻的键倾向于具有相似的重要性水平。
为了实现区块选择,我们首先将键和值序列划分为选择区块。为了识别对注意力计算至关重要的区块,我们需要为每个区块分配重要性分数。
重要性分数
压缩 token 的注意力计算会产生中间注意力分数,可以通过利用这些分数来推导选择区块的重要性分数:
ptcmp=Softmax(qtTK~tcmp)
\mathbf{p}_{t}^{\text{cmp}}=\text{Softmax}(\mathbf{q}_{t}^\mathrm T\tilde{K}_t^{\text{cmp}})
ptcmp=Softmax(qtTK~tcmp)
其中 ptcmp∈R⌊t−ld⌋\mathbf{p}_{t}^{\text{cmp}} \in \mathbb{R}^{\left\lfloor \frac{t-l}{d} \right\rfloor}ptcmp∈R⌊dt−l⌋ 表示与 qt\mathbf{q}_tqt 与压缩键 K~tcmp\tilde{K}_t^{\text{cmp}}K~tcmp 之间的注意力分数。令 l′l'l′ 表示选择块的大小,
- 当压缩块和选择块共享相同的分块方案,即 l′=l=dl' = l = dl′=l=d 时,选择块的重要性分数 ptslc\mathbf{p}_{t}^{\text{slc}}ptslc 可以直接通过 ptslc=ptcmp\mathbf{p}_{t}^{\text{slc}}=\mathbf{p}_{t}^{\text{cmp}}ptslc=ptcmp 得到。
- 当分块方案不同时,根据选择块的空间关系推导其重要性分数。给定 d∣ld \mid ld∣l 且 d∣l′d \mid l'd∣l′:
ptslc[j]=∑m=0l′d−1∑n=0ld−1ptcmp[l′dj+m+n] \mathbf{p}^{\text{slc}}_t[j] = \sum_{m=0}^{\frac{l'}{d}-1}\sum_{n=0}^{\frac{l}{d}-1} \mathbf{p}^{\text{cmp}}_t[\frac{l'}{d}j + m + n] ptslc[j]=m=0∑dl′−1n=0∑dl−1ptcmp[dl′j+m+n]
其中 [⋅][·][⋅] 表示用于访问向量元素的索引运算符。
- 对于采用 GQA 或 MQA 的模型,其中键值缓存跨查询头共享,在解码过程中需要确保这些头的一致性块选择,以最小化键值缓存的加载。组内各头之间的共享重要性得分:
ptslc′=∑h=1Hptslc′(h) \mathbf{p}_t^{\text{slc}'} = \sum_{h=1}^{H} \mathbf{p}_t^{\text{slc}'(h)} ptslc′=h=1∑Hptslc′(h)
其中,上标中的 (h)(h)(h) 表示头部索引,HHH 是每个组中查询头部的数量,这种聚合确保了同一组内各头部之间的块选择一致性。
Top-n 区块选择
在获得选择区块的重要性分数后,保留按重要性分数排名前 n 的稀疏块内的 token:
It={i∣rank(ptslc′[i])≤n}K~tslc=Cat[{kil′+1:(i+1)l′∣i∈It}]
I_t = \{i \mid \text{rank}(\mathbf{p}^{\text{slc}'}_t[i]) \le n\}\\
\tilde{K}^{\text{slc}}_t = \text{Cat}[\{\mathbf{k}_{il'+1:(i+1)l'} \mid i \in I_t\}]
It={i∣rank(ptslc′[i])≤n}K~tslc=Cat[{kil′+1:(i+1)l′∣i∈It}]
其中,rank(⋅)\text{rank}(·)rank(⋅) 表示按降序排列的排名位置,ItI_tIt 表示选择区块的索引集,Cat\text{Cat}Cat 表示拼接操作。K~tslc∈Rdk×nl′\tilde{K}^{\text{slc}}_t\in \mathbb{R}^{d_k \times nl'}K~tslc∈Rdk×nl′ 是由压缩键组成的张量。**细粒度值同理。**随后,所选的键和值将与 qt\mathbf{q}_{t}qt 参与注意力计算。
3.3.3 滑动窗口
在注意力机制中,局部模式通常适应速度较快,并可能主导学习过程,从而阻碍模型从压缩和选择 token 中有效学习。为解决这一问题,我们引入了一个专门的滑动窗口分支,该分支明确处理局部上下文,使其他分支(压缩和选择)能够专注于学习各自的功能,而不受局部模式的干扰。
具体而言,维护一个窗口 www 内的最近 token:
- K~twin=kt−w:t\tilde{K}_t^{\text{win}} = \mathbf{k}_{t-w:t}K~twin=kt−w:t
- V~twin=vt−w:t\tilde{V}_t^{\text{win}} = \mathbf{v}_{t-w:t}V~twin=vt−w:t
分别将不同信息源的注意力计算**(压缩 token、选择 token 和滑动窗口)隔离到单独的分支中。这些分支的输出通过一个学习到的门控机制进行聚合。为了进一步防止注意力分支之间的捷径学习,同时仅引入微小的计算开销,我们为三个分支提供独立的键和值**。这种架构设计通过防止局部和长距离模式识别之间的梯度干扰,实现了稳定的训练,同时引入了最小的开销。在获得所有三类键和值 (K~tcmp,V~tcmp;K~tslc,V~tslc;K~twin,V~twin)(\tilde{K}_t^{\text{cmp}},\tilde{V}_t^{\text{cmp}}; \tilde{K}_t^{\text{slc}},\tilde{V}_t^{\text{slc}};\tilde{K}_t^{\text{win}},\tilde{V}_t^{\text{win}})(K~tcmp,V~tcmp;K~tslc,V~tslc;K~twin,V~twin) 之后,计算最终的注意力输出。
3.4 内核设计
为了在训练和预填充过程中实现类似于 FlashAttention 的速度提升,我们基于 Triton 实现了硬件对齐的稀疏注意力内核。鉴于多头注意力(MHA)在内存密集型且解码效率低下,我们专注于采用共享键值缓存的架构,如 GQA 和 MQA。
虽然压缩和滑动窗口注意力计算与现有的 FlashAttention-2 内核兼容,我们仍引入了专门设计的稀疏选择注意力内核。如果我们沿用 FlashAttention 的策略,将时间连续的查询块加载到SRAM中,这会导致不高效的内存访问,因为一个块内的查询可能需要不连续的键值块。为了解决这一问题,我们的关键优化在于不同的查询分组策略:对于查询序列中的每个位置,我们将同一GQA组内的所有查询头(它们共享相同的稀疏键值块)加载到SRAM中。

内核架构具有以下关键特征:
- **群体中心的数据加载:**对于每个内循环,加载群体中位置为 ttt 的所有头部的查询 Q∈R[h,dk]Q \in \mathbb{R}^{[h,d_k]}Q∈R[h,dk] 以及它们共享的稀疏键/值块索引 ItI_tIt。
- **共享键值读取:**在内层循环中,按索引 ItI_tIt 顺序加载连续的键/值块 K∈R[Bk,dk]K\in\mathbb{R}^{[B_k,d_k]}K∈R[Bk,dk] 和 V∈R[Bk,dv]V\in \mathbb{R}^{[B_k,d_v]}V∈R[Bk,dv] 到 SRAM 中以最小化内存加载,其中 BkB_kBk 是满足 Bk∣l′B_k | l'Bk∣l′ 的核块大小。
- **在外层网格循环:**由于不同查询块的内层循环长度(与选定的块数 nnn 成正比)几乎保持不变,我们将查询/输出循环放入 Triton 的网格调度器中以简化并优化内核。
该设计通过(1)组内共享消除冗余的键值传输(2)在GPU流式多处理器之间平衡计算工作负载,实现了接近最优的算术强度。
四、实验
从三个角度评估 NSA:
- 通用基准性能
- 长上下文基准性能
- 链式推理性能
并将其与全注意机制基线和最先进的稀疏注意力方法进行比较。
4.1 预训练设置
遵循当前最先进的大语言模型(LLM)的常见做法,实验采用了一种结合分组查询注意(GQA)和专家混合(MoE)的骨干网络:
- 总参数量为270亿,其中活跃参数量为30亿。
- 模型包含 30 层,隐藏维度为 2560。
- 对于 GQA,将组的数量设置为 4,总共有 64 个注意力头。
- 对于每个头,查询、键和值的隐藏维度分别配置为 dq=dk=192d_q = d_k = 192dq=dk=192 和 dv=128d_v = 128dv=128。
- 对于 MoE,采用了 DeepSeekMoE 结构,包括 72 个路由专家和 2 个共享专家,并将 top-k 专家设置为6。为了确保训练稳定性,在第一层的MoE被 SwiGLU 形式的多层感知机(MLP)所替代。
提出的架构在计算成本和模型性能之间实现了有效的权衡。对于NSA,我们设置:
- 压缩区块大小 l=32l = 32l=32
- 滑动步长 d=16d = 16d=16
- 选择区块大小 l′=64l' = 64l′=64
- 选择区块数量 n=16n = 16n=16(包括固定激活的1个初始块和2个局部块)
- 滑动窗口大小 w=512w = 512w=512
全注意力模型和稀疏注意力模型均在长度为 8k 的 270B 词元上进行预训练,随后在长度为 32k 的词元上进行持续训练和监督微调,使用 YaRN 实现长上下文适应。两个模型均训练至完全收敛,以确保公平比较。

NSA 和 全注意力模型基线的预训练损失曲线均表现出稳定且平滑的下降趋势,NSA在整个过程中持续优于全注意力模型。
4.2 基线方法
除了与全注意机制进行比较外,我们还评估了几种最新的推理阶段稀疏注意力方法:H2O、infLLM、Quest 以及 Exact-Top(首先计算全注意力分数,然后选择与每个查询对应的前 n 个最高分数的关键位置,并在此基础上计算注意力)。这些方法涵盖了多种稀疏注意机制,包括 KV 缓存淘汰、查询感知选择和精确的 Top-n 稀疏选择。
- **通用基准性能:**大多数样本的长度都在稀疏注意力基线的局部上下文窗口范围内,因此这些方法在本质上等同于全注意力。因此,在这种设置下,我们仅呈现 NSA 与全注意力基线的比较结果。
- **长上下文基准性能:**对所有基线方法进行了比较,并将所有稀疏注意力方法的稀疏性设置为相同,以确保比较的公平性。
- **思维链推理性能:**仅将比较限定于全注意力基线,因为稀疏注意力的基线方法不支持训练。
4.3 性能比较
4.3.1 通用评估

4.3.2 长上下文评估


4.3.3 思维链推理评估

五、效率分析
评估 NSA 在 8 块 A100 GPU 系统上的计算效率,与全注意机制进行了对比。
-
模型配置: GQA 组 g=4g = 4g=4,每组头数 h=16h = 16h=16,查询/键维度 dk=192d_k = 192dk=192,值维度 dv=128d_v = 128dv=128。
-
NSA: 压缩区块大小 l=32l = 32l=32,滑动步长 d=16d = 16d=16,选择区块大小 l′=64l' = 64l′=64,选择区块数量 n=16n = 16n=16,滑动窗口大小 w=512w = 512w=512
5.1 训练速度

5.2 解码速度

六、讨论
6.1 替代性 token 选择策略面临的挑战
在设计NSA之前,我们探索了将现有的稀疏注意力方法应用于训练阶段的可能性。然而,这些尝试遇到了各种挑战,促使我们设计了一种不同的稀疏注意力架构:
**基于键聚类的策略。**ClusterKV 等基于聚类的策略将来自同一聚类的键和值存储在连续的内存区域中,虽然在训练和推理方面理论上是可行的,但它们面临三个显著挑战:
- 由动态聚类机制引入的重大计算开销;
- 由于跨聚类不平衡导致的操作优化困难,特别是在混合专家系统(MoE)中,偏差的专家并行执行时间(EP)导致持续的负载不平衡;
- 由于需要强制进行周期性重新聚类和块顺序训练协议而产生的实现约束。
**其他分块选择策略。**Quest 和 InfLLM 等与 NSA 不同的分块键和值选择策略,依赖于为每个分块计算重要性得分,并基于其与 qtq_tqt 的相似度选择前 nnn 个分块。然而,现有方法面临两个关键问题:
- 由于选择操作是非可微的,基于神经网络的重要性得分计算依赖于辅助损失,这增加了操作员的开销并通常会降低模型性能;
- 基于启发式无参数的重要性得分计算策略会导致召回率低,从而导致性能不佳。
我们使用具有相似架构的3B参数模型评估这两种方法,并将它们的损失曲线与 NSA 和全注意机制进行比较。
-
对于基于辅助损失的选择方法,为每个分块引入额外的查询和代表性的键以估计分块的重要性得分,这些得分由原始查询和每个分块内键的平均注意力得分监督。
-
对于基于启发式的无参数选择方法,遵循Quest的策略,直接使用查询与键分块的坐标最小-最大值的乘积进行选择,而不引入额外的参数。
我们还探索了一种冷启动训练方法,即在初始的1000步中使用全注意机制,之后过渡到启发式的分块选择。

6.2 可视化
为了探索 transformer 注意力分布中的潜在模式,并为设计寻找灵感,我们可视化了预训练的 27B 全注意力模型的注意力图。

可视化结果显示,注意力分数倾向于表现出块状聚类的特征,相邻的键往往具有相似的注意力分数。这一观察结果启发了我们设计 NSA,表明基于空间连续性选择键块可能是一种有前景的方法。
块状聚类现象表明,序列中相邻的标记可能与查询标记共享某些语义关系,尽管这些关系的具体性质仍需进一步研究。这一观察结果促使我们探索一种基于连续标记块而非单个标记的操作稀疏注意力机制,旨在提高计算效率并保留高注意力模式。
为了探索 transformer 注意力分布中的潜在模式,并为设计寻找灵感,我们可视化了预训练的 27B 全注意力模型的注意力图。
[外链图片转存中…(img-IyzSznIQ-1739953314036)]
可视化结果显示,注意力分数倾向于表现出块状聚类的特征,相邻的键往往具有相似的注意力分数。这一观察结果启发了我们设计 NSA,表明基于空间连续性选择键块可能是一种有前景的方法。
块状聚类现象表明,序列中相邻的标记可能与查询标记共享某些语义关系,尽管这些关系的具体性质仍需进一步研究。这一观察结果促使我们探索一种基于连续标记块而非单个标记的操作稀疏注意力机制,旨在提高计算效率并保留高注意力模式。
更多推荐



所有评论(0)