大语言模型研究16——注意力机制优化之稀疏注意

作者: 引线小白-本文永久链接:https://www.limoncc.com/post/1540fb65f18f359a/
知识共享许可协议: 本博客采用署名-非商业-禁止演绎4.0国际许可证

一、什么是DSA

1.1、核心技巧

DeepSeek DSA = DeepSeek Sparse Attention[^1],是 DeepSeek 团队提出的一种“细粒度稀疏注意力机制”,专门用来在长上下文场景下大幅降低注意力计算和显存开销,同时尽量保持模型性能不变。

标准 Transformer 自注意力的复杂度是:

  • 计算量:$O(L^2)$(序列长度 L)
  • KV Cache 显存:$O(L)$

当上下文从 4K → 128K → 1M token,计算和显存会爆炸式增长,长上下文非常贵。 DSA 的目标就是:让长上下文训练/推理又快又省,同时不把模型搞笨。DSA 的核心思路(通俗版)一句话概括: 不是所有 token 都需要“一对一精算注意力”,DSA 用一个“闪电索引器(Lightning Indexer)”先选出少量关键token,只对这些做精细注意力,其它的或被跳过、或被粗略处理。可以理解为:

  1. 先用“轻量筛选器”给每个 query 选出最相关的 top-K 个 key(token 位置);
  2. 只在这些 top-K 位置上做正常的注意力计算;
  3. 再把稀疏注意力的结果和主路径特征融合。

这样,注意力就从“全连接”变成“稀疏连接”,计算量从 $O(L^2)$ 降到接近 $O(L \cdot K)$,K ≪ L。

1.2、DSA核心数学公式

回顾标准自注意力, 给定输入序列 $\bm{X} \in \mathbb{R}^{L \times d}$($L$ 为序列长度,$d$ 为隐藏维度),标准自注意力计算为:

$$\begin{align}
\bm{Q} = \bm{X} \bm{W}_Q, \quad
\bm{K} = \bm{X} \bm{W}_K, \quad
\bm{V} = \bm{X} \bm{W}_V
\end{align}$$

$$\begin{align}
\text{Attention}(\bm{Q}, \bm{K}, \bm{V}) = \text{softmax}\left( \frac{\bm{Q} \bm{K}^\top}{\sqrt{d_k}} \right) \bm{V}
\end{align}$$

其中 $d_k$ 是 Key 的维度,计算复杂度为 $O(L^2 d)$。

DSA 的核心思想:对每个 Query,只选择与其最相关的 $K$ 个 Key($K \ll L$)参与注意力计算。选择过程通过一个轻量级的 Indexer 完成。

步骤 1:Indexer 计算 Query 和 Key

Indexer 使用独立的投影矩阵,从输入 $\bm{X}$ 和 MLA 压缩后的 Query 表示 $\bm{q}_r$ 生成索引用的 Query 和 Key:

$$\begin{align}
\bm{Q}_{\text{idx}} = \text{RoPE}\left( \bm{q}_r \bm{W}_{Q_{\text{idx}}} \right), \quad
\bm{K}_{\text{idx}} = \text{RoPE}\left( \bm{X} \bm{W}_{K_{\text{idx}}} \right)
\end{align}$$

$\bm{q}_r \in \mathbb{R}^{L \times d_q}$:MLA 压缩后的 Query 表示($d_q \ll d$)
$\bm{W}_{Q_{\text{idx}}} \in \mathbb{R}^{d_q \times d_{\text{idx}}}$, $\bm{W}_{K_{\text{idx}}} \in \mathbb{R}^{d \times d_{\text{idx}}}$:Indexer 的投影矩阵
$\text{RoPE}$:旋转位置编码
$d_{\text{idx}}$:Indexer 的维度(通常远小于 $d$)

步骤 2:计算相关性分数并选择 Top-K 位置

对每个 Query 位置 $i$,计算其与所有 Key 位置的相关性分数,并选出分数最高的 $K$ 个位置:
$$\begin{align}
\bm{S}_i &= \bm{Q}_{\text{idx},i} \bm{K}_{\text{idx}}^\top \in \mathbb{R}^{1 \times L}\\
\mathcal{I}_i &= \text{TopK}(\bm{S}_i, K) \subset \{1, 2, \dots, L\}, \quad |\mathcal{I}_i| = K
\end{align}$$

其中:
$\mathcal{I}_i$:位置 $i$ 的稀疏索引集合
实际实现中,该步骤通过优化的 FP8 Kernel(如 fp8_index)高效完成

步骤 3:构造稀疏掩码

根据 Top-K 索引构造二值掩码:

$$\begin{align}
\bm{M}_{i,j} =
\begin{cases}
1, & \text{if } j \in \mathcal{I}_i \\
0, & \text{otherwise}
\end{cases}
\end{align}$$

步骤 4:稀疏注意力计算

在主注意力路径(MLA)中,只对掩码选中的位置计算注意力:
$$\begin{align}
\bm{Q} = \bm{X} \bm{W}_Q, \quad
\bm{K} = \bm{X} \bm{W}_K, \quad
\bm{V} = \bm{X} \bm{W}_V
\end{align}$$
对于每个 Query 位置 $i$:

$$\begin{align}
\bm{A}_i = \text{softmax}\left( \frac{\bm{Q}_i \bm{K}_{\mathcal{I}_i}^\top}{\sqrt{d_k}} \right) \bm{V}_{\mathcal{I}_i}
\end{align}$$

  • $ \bm{K}_{\mathcal{I}_i}, \bm{V}_{\mathcal{I}_i} $:仅索引集合 $ \mathcal{I}_i $ 中位置的 Key 和 Value

  • 注意:在 DeepSeek 的 MLA 中,$\bm{K}$ 和 $\bm{V}$是从压缩表示恢复的,但这里为清晰起见省略了压缩步骤。最终输出为所有 $\bm{A}_i $ 拼接的结果。

复杂度对比

方法 计算复杂度 KV Cache 显存
标准注意力 $O(L^2 d)$ $O(L d)$
DSA 稀疏注意力 $O(L K d)$ $O(L d)$ 但计算时只需访问 $K$ 个位置

由于 $K \ll L$(例如 $L=128K$, $K=2048$ ),计算量从二次降为线性级别。

二、数值例子

为直观理解,我们构造一个极简的数值例子,假设:

  • 序列长度 $L=4$,隐藏维度 $d=2$
  • Top-K 参数 $K=1$(每个 Query 只选 1 个最相关的 Key)
  • 忽略 RoPE、MLA 压缩等细节,聚焦核心选择逻辑

1.输入序列

$$\begin{align}
\bm{X} = \begin{bmatrix}
1 & 2 \\
3 & 4 \\
5 & 6 \\
7 & 8
\end{bmatrix} \in \mathbb{R}^{4 \times 2}
\end{align}$$

2.Indexer 投影(简化)

假设 Indexer 的投影矩阵为:

$$\begin{align}
\bm{W}_{Q_{\text{idx}}} = \begin{bmatrix} 1 & 0 \ 0 & 1 \end{bmatrix}, \quad
\bm{W}_{K_{\text{idx}}} = \begin{bmatrix} 1 & 1 \ 1 & -1 \end{bmatrix}
\end{align}$$

计算 Indexer 的 Query 和 Key(忽略 RoPE):
$$\begin{align}
\bm{Q}_{\text{idx}} = \bm{X} \bm{W}_{Q_{\text{idx}}} = \bm{X} = \begin{bmatrix}
1 & 2 \\
3 & 4 \\
5 & 6 \\
7 & 8
\end{bmatrix}
\end{align}$$

$$\begin{align}
\bm{K}_{\text{idx}} = \bm{X} \bm{W}_{K_{\text{idx}}} = \begin{bmatrix}
1\cdot1+2\cdot1 & 1\cdot1+2\cdot(-1) \\
3\cdot1+4\cdot1 & 3\cdot1+4\cdot(-1) \\
5\cdot1+6\cdot1 & 5\cdot1+6\cdot(-1) \\
7\cdot1+8\cdot1 & 7\cdot1+8\cdot(-1)
\end{bmatrix} = \begin{bmatrix}
3 & -1 \\
7 & -1 \\
11 & -1 \\
15 & -1
\end{bmatrix}
\end{align}$$

3.计算相关性分数并选择 Top-1

计算相关性矩阵 $\bm{S} = \bm{Q}_{\text{idx}} \bm{K}_{\text{idx}}^\T$:

$$\begin{align}
\bm{S} = \begin{bmatrix}
1\cdot3 + 2\cdot(-1) & 1\cdot7 + 2\cdot(-1) & 1\cdot11 + 2\cdot(-1) & 1\cdot15 + 2\cdot(-1) \\
3\cdot3 + 4\cdot(-1) & 3\cdot7 + 4\cdot(-1) & 3\cdot11 + 4\cdot(-1) & 3\cdot15 + 4\cdot(-1) \\
5\cdot3 + 6\cdot(-1) & 5\cdot7 + 6\cdot(-1) & 5\cdot11 + 6\cdot(-1) & 5\cdot15 + 6\cdot(-1) \\
7\cdot3 + 8\cdot(-1) & 7\cdot7 + 8\cdot(-1) & 7\cdot11 + 8\cdot(-1) & 7\cdot15 + 8\cdot(-1)
\end{bmatrix} = \begin{bmatrix}
1 & 5 & 9 & 13 \\
5 & 17 & 29 & 41 \\
9 & 29 & 49 & 69 \\
13 & 41 & 69 & 97
\end{bmatrix}
\end{align}$$

对每行(每个 Query)选择最大值对应的索引(Top-1):

Query 0:最大值 13 → 索引 3 → $ \mathcal{I}_0 = \{3\} $
Query 1:最大值 41 → 索引 3 → $ \mathcal{I}_1 = \{3\} $
Query 2:最大值 69 → 索引 3 → $ \mathcal{I}_2 = \{3\} $
Query 3:最大值 97 → 索引 3 → $ \mathcal{I}_3 = \{3\} $

在这个例子中,所有 Query 都只关注位置 3(最后一个 token)。

4.稀疏注意力计算

假设主注意力路径的投影矩阵均为单位矩阵(简化),即 $\bm{Q}=\bm{K}=\bm{V}=\bm{X}$。
以 Query 0 为例:

1、$\bm{Q}_0 = [1, 2]$
2、只对 Key 位置 3 计算:$\bm{K}_3 = [7, 8]$, $\bm{V}_3 = [7, 8]$
3、注意力分数:$\bm{Q}_0 \bm{K}_3^\top = 1\cdot7 + 2\cdot8 = 23$
4、Softmax(只有一个选项):$\text{softmax}(23/\sqrt{2}) \approx 1$
5、输出:$\bm{A}_0 = 1 \cdot [7, 8] = [7, 8]$

同理:
Query 1:$\bm{Q}_1 = [3, 4]$, 分数 = 53 → 输出 $[7, 8]$
Query 2:$\bm{Q}_2 = [5, 6]$, 分数 = 83 → 输出 $[7, 8]$
Query 3:$\bm{Q}_3 = [7, 8]$, 分数 = 113 → 输出 $[7, 8]$
最终输出序列所有位置均为 $[7, 8]$。

该例子中,Indexer 的投影使得所有 Query 都最关注最后一个 token,因此稀疏注意力只聚合了该位置的信息。

实际中,$K$ 会更大(如 2048),且 Indexer 会学习到更丰富的模式(如局部性、语义相关性),不会出现所有 Query 只关注一个 token 的极端情况。

DSA 的关键优势:通过 Indexer 动态选择少量关键位置,将计算量从 $O(L^2)$ 降至 $O(LK)$,同时保持模型性能。

三、DSA梯度问题

3.1、直通估计器

直通估计器(Straight-Through Estimator,简称 STE)是深度学习中用于解决“离散操作不可导”问题的经典技巧。它核心思想是“欺骗”梯度计算:在前向传播时做硬离散化(保证推理效果),在反向传播时假装前向做的是恒等映射(保证梯度流动)。

3.1.1、为什么需要 STE?

深度学习依赖反向传播算法,这要求网络中的每一步操作必须是可导的(或者至少是次可导的)。然而,很多操作本质上是离散的,比如:

1.二值化:$x \in \mathbb{R} \to y \in \{-1, +1\}$。
2.量化:将浮点数转为定点整数。
3.采样:从离散分布中取出一个样本。
4.Top-K/Argmax:选出最大的那个数,输出是一个不可导的索引或硬掩码(DSA 中的场景)。

这些函数的导数在数学上要么是 0(对于非边界点),要么未定义。如果直接用它们,梯度会断流,模型无法训练。
STE 的解决方案非常粗暴但有效:既然离散函数的导数没法算,那我们在反向传播时直接“无视”这个离散操作,把后层的梯度直接传给前层。也就是把离散操作当作“直通管道”来传梯度。

假设有一个离散化函数 $\mathrm{operator}(\cdot)$(比如 Sign、Round 或 Top-K),它将连续输入 $x$ 映射为离散输出 $y$:

$$\begin{align}
y=\mathrm{operator}(x)
\end{align}$$

在前向传播中,我们老老实实地用 $y$。但在反向传播计算梯度时,STE 假设 $y$ 和 $x$ 之间存在如下关系: $y \approx x$。因此,梯度满足 $\displaystyle \frac{\partial y}{\partial x} \approx 1$, 根据链式法则,损失函数 $L$ 对 $x$ 的梯度为:

$$\begin{align}
\frac{\partial \ell}{\partial x} = \frac{\partial \ell}{\partial y} \cdot \frac{\partial y}{\partial x} \approx \frac{\partial \ell}{\partial y} \cdot 1 = \frac{\partial \ell}{\partial y}
\end{align}$$

这就是“直通”的含义:后层的梯度 $\frac{\partial \ell}{\partial y}$ 直接穿透离散操作,原封不动地变成了前层的梯度 $\frac{\partial \ell}{\partial x}$。

3.1.2、经典例子:二值神经网络

二值神经网络是 STE 最原始、最经典的应用场景:

要把权重 $w_r$(实数)压缩成二值权重 $w_b \in \{-1, +1\}$ 以节省存储和加速计算。前向传播(硬前向)
使用 Sign 函数进行硬离散化:

$$\begin{align}
w_b = \text{sign}(w_r) = \begin{cases}
+1, & w_r \ge 0 \\
-1, & w_r < 0
\end{cases}
\end{align}$$

前向计算只用 $w_b$,模型在推理时是纯二值的。反向传播(STE)中Sign 函数的导数几乎处处为 0(除了 0 点处无定义)。如果直接求导,梯度全为 0,训练终止。应用 STE,我们忽略 Sign 函数的存在,假设 $w_b \approx w_r$:

$$\begin{align}
\frac{\partial L}{\partial w_r}
= \frac{\partial L}{\partial w_b} \cdot \underbrace{\frac{\partial \text{sign}(w_r)}{\partial w_r}}_{\approx 1 \text{ (via STE)}}
\approx \frac{\partial L}{\partial w_b}
\end{align}$$

假设 Loss 告诉我们需要增加 $w_b$ 的值。虽然 $w_b$ 只能是 $+1$ 或 $-1$,无法微调,但 STE 把这个“增加”的意图传给了背后的实数潜影 $w_r$。 $w_r$ 收到梯度后会变大,比如从 $0.1$ 变成 $0.5$。由于 $w_r > 0$,前向时 $w_b$ 依然是 $+1$。但如果 $w_r$ 继续变大,$w_b$ 的状态就更稳固了;反之,如果 $w_r$ 变小甚至变成负数,$w_b$ 就会从 $+1$ 翻转为 $-1$。 STE 让原本不可调的二值权重,通过背后的实数潜影实现了“翻牌”的可能。

3.2、DSA的反向传播

稀疏注意力(乃至所有包含硬路由/Top-K机制的网络,如MoE)的有一个核心痛点:Top-K 操作在数学上是不可导的(或者更准确地说,其梯度几乎处处为0),这会切断反向传播的链条。如果直接用 Top-K 生成 0/1 掩码,Indexer 的参数($\bm{W}_{Q_{idx}}$ 和 $\bm{W}_{K_{idx}}$)将无法获得梯度,也就无法学习到“到底哪些 Token 是重要的”。为了解决这个问题,DeepSeek 在 DSA 的训练中采用了业界标准的直通估计器,并结合了软松弛技术。具体来说,它是这样做反向梯度的:

假设 Indexer 计算出的相关性分数为 $\bm{S} = \bm{Q}_{idx} \bm{K}_{idx}^\T$。Top-K 操作生成的硬掩码为 $\bm{M}_{hard}$:

$$\begin{align}
M_{hard, i} = \begin{cases}
1, & \text{if } S_i \in \text{Top-K} \\
0, & \text{otherwise}
\end{cases}
\end{align}$$

最终输出 $\bm{Y} = f(\bm{X} \odot \bm{M}_{hard})$。
求导时:

$$\begin{align}
\frac{\partial \bm{Y}}{\partial S_i} = \frac{\partial \bm{Y}}{\partial M_{hard, i}} \cdot \frac{\partial M_{hard, i}}{\partial S_i}
\end{align}$$

因为 $M_{hard}$ 是阶跃的0/1输出,$\frac{\partial M_{hard, i}}{\partial S_i}$ 在绝大部分地方都是 0。未被选中的 Token 永远得不到梯度,已被选中的 Token 也无法根据分数微调自己的重要性,Indexer 陷入了“盲人摸象”的死局。DeepSeek DSA 训练时的核心逻辑是:前向传播为了保证推理一致性和计算效率,使用硬 Top-K;反向传播为了让参数能学习,假装前向做的是平滑的 Softmax/Sigmoid。这就是 Straight-Through Estimator (STE) 的思想。

1.前向传播

正常计算 $S$,并应用硬 Top-K 掩码:

$$\begin{align}
\bm{M}_{hard} = \text{TopKMask}(\bm{S})
\end{align}$$

主注意力路径使用 $\bm{M}_{hard}$ 进行计算,这部分是确定的、稀疏的。

2.反向传播

在反向求导时,我们“偷梁换柱”,用一个可导的连续函数 $\bm{M}_{soft}$ 的导数,去替代 $\bm{M}_{hard}$ 的导数。最常见的替代方案是 Sigmoid 松弛 或 Softmax 松弛。常见有如下思路:

方案 A:Sigmoid 松弛

假设我们用 Sigmoid 函数把分数映射到 (0,1) 作为软掩码:$\bm{M}_{soft} = \sigma(\bm{S} / \tau)$ ($\tau$ 是温度系数,越小越接近0/1)。
反向传播时,梯度按 Sigmoid 的导数传递:

$$\begin{align}
\frac{\partial \ell}{\partial S_i} \approx \frac{\partial L}{\partial M_{hard, i}} \cdot \sigma’(S_i / \tau)
\end{align}$$

  • 对于被 Top-K 选中的 Token($M_{hard}=1$):它获得了来自主注意力的梯度,乘以 Sigmoid 导数,鼓励它保持高分。
  • 对于未被选中的 Token($M_{hard}=0$):虽然主路径的梯度 $\frac{\partial L}{\partial M_{hard, i}}$ 传不过来(因为乘了0),但 STE 的变体通常会强制给未选中但分数接近阈值的 Token 一个小的梯度(称为梯度泄露),或者通过 Sigmoid 导数在分数接近边界时提供微弱的推力,告诉它“你差一点就进来了,再努力一点”。

方案 B:Gumbel-Softmax(更常用于离散路由)

Gumbel-Softmax 是一种更精细的离散化近似技术。
前向传播时,加入随机噪声(Gumbel 噪声)并取 argmax(等效于 Top-1 的随机版本),反向传播时,使用 Softmax 的梯度:

$$\begin{align}
\bm{M}_{hard} = \text{OneHot}(\text{argmax}(\bm{S} + \text{Gumbel})) \quad \text{(前向)}
\end{align}$$

$$\begin{align}
\frac{\partial \bm{M}_{hard}}{\partial S} \approx \frac{\partial \text{Softmax}((\bm{S} + \text{Gumbel})/\tau)}{\partial \bm{S}} \quad \text{(反向)}
\end{align}$$

对于 Top-K,通常是执行 K 次 Gumbel-Softmax 或者使用专门的 Top-K Gumbel 松弛。

3.3、DSA梯度问题特殊处理

DeepSeek 在 DSA 的实现中,为了让 STE 和 FP8 训练结合,做了如下底层定制:

1、FP8 前向,BF16 反向:

  • 前向传播时,Indexer 的 $\bm{Q}_{idx}$ 和 $\bm{K}_{idx}$ 乘法是在 FP8 下极致加速的,算出来的分数 $\bm{S}$ 也是 FP8。
  • 但是,FP8 的动态范围太小,无法精确表达 STE 需要的 Sigmoid/Softmax 导数。
  • 因此,在自定义的 fp8_index CUDA Kernel 中,前向用 FP8 算 Top-K,但在反向传播时,会基于保存的 BF16 精度的 $\bm{S}$ 或 $\bm{Q}_{idx}, \bm{K}_{idx}$ 来计算 STE 梯度。

2、梯度缩放与裁剪:

  • STE 估计出来的梯度往往方差很大,容易导致训练不稳定。DeepSeek 在反向传播时,会对 $\displaystyle \frac{\partial \ell}{\partial S}$ 进行缩放和裁剪,确保 Indexer 的权重更新平滑。

3、辅助负载均衡损失(Auxiliary Loss,可选):

  • 类似于 MoE 中的做法,如果 Top-K 总是倾向于选择序列中某几个“明星 Token”,会导致路由崩塌。
  • 虽然注意力机制不像 MoE 那样严格需要均衡,但为了防止 Indexer 陷入局部最优(比如永远只看最近的几个 Token),有时会附加一个轻量级的辅助损失,鼓励 Indexer 的选择在位置上具有一定分散性:

$$\begin{align}
\ell_{aux} = \lambda \cdot \text{Var}(\text{Count}_{\text{selected}})
\end{align}$$
这个辅助损失的梯度是可以直接传给 $S$ 的,不需要经过 Top-K,从而保证了梯度的基本流动。

3.4、为什么不是 Gumbel-Softmax

Gumbel-Softmax 虽然是处理离散采样的经典方法,但在 LLM 的大规模训练中,它有两个致命缺陷,DeepSeek 不可能采用:

1、注入随机噪声导致训练崩塌:

Gumbel-Softmax 的核心是在前向分数中加入 Gumbel 噪声来逼近 argmax。但在像 DeepSeek-V3 这样拥有 671B 参数、使用 FP8 精度进行大规模训练的模型中,数值区间本身就极小(FP8 的动态范围很窄)。注入随机噪声极易导致梯度方差爆炸,直接摧毁 FP8 训练的稳定性。

2、计算开销极大:

Gumbel-Softmax 需要对整个词表长度 $L$ 计算完整的 Softmax(即使分块计算也很昂贵),这完全违背了 DSA “用极低成本做初筛”的初衷。Indexer 的设计初衷就是用几行代码、极小的算力把 Top-K 选出来,Gumbel-Softmax 的计算代价太昂贵了。

3.5、 直观理解

你可以把 DSA 的训练想象成一场“带保底机制的选拔赛”:
1.前向(选拔):严格按分数选 Top-K,没选上的直接淘汰,不让它们参与最终的计算(硬掩码)。
2.反向(复盘):

  • 选上的选手,如果对最终结果有正面贡献,会得到正向反馈(梯度),鼓励它们保持高分;
  • 没选上的选手,虽然没参与计算,但 STE 机制相当于给那些“分数只差一点点就能进 Top-K”的选手发了短信:“你其实很有潜力,下次再把分数考高一点就能入选了”。
  1. 通过成千上万次的“选拔-复盘”,Indexer 的参数 $\bm{W}_{Q_{idx}}$ 和 $\bm{W}_{K_{idx}}$
3.6、代码层面的工程实现

在实际的 PyTorch / Triton 代码实现中,你通常看不到显式写出 Sigmoid 松弛公式,而是通过一种更优雅的 截断反向传播 来实现 STE:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
# 前向:正常算分,取 Top-K 硬掩码
scores = q_idx @ k_idx.T
threshold = torch.kthscore(scores, k=K)
mask_hard = (scores >= threshold).float()
# 反向:自定义梯度
class DSATopK(torch.autograd.Function):
@staticmethod
def forward(ctx, scores, K):
# ... 算 Top-K 并返回 mask_hard
ctx.save_for_backward(scores, threshold)
return mask_hard
@staticmethod
def backward(ctx, grad_output):
scores, threshold = ctx.saved_tensors
# 只有被选中的位置,梯度原样透传(直通)
grad_scores = grad_output * (scores >= threshold).float()

# 可选:为了防止边界震荡,给边界附近的负样本一点微小梯度
# distance = scores - threshold
# grad_scores += grad_output * sigmoid_derivative(distance) * (scores < threshold).float()

return grad_scores, None

在 DeepSeek 底层用 Triton/Hack 实现的 fp8_index Kernel 中,这个逻辑被写成了 前向 FP8 算 Top-K,反向用 BF16 算 STE 透传 的融合算子。

四、工程实现:它是怎么做到“又快又准”的?

如果只看数学公式,算 $\bm{Q}_{idx} \times \bm{K}_{idx}^T$ 依然是 $O(L^2)$。Indexer 之所以能落地,是因为 DeepSeek 团队在底层硬件层面做了极致的定制:

1、Hadamard 旋转与 FP8 量化:在算内积之前,对 $\bm{Q}_{idx}$ 和 $\bm{K}_{idx}$ 做 Hadamard 旋转,使得向量各个维度的分布变得均匀,从而可以安全地使用 FP8(8位浮点数) 进行极高效率的矩阵乘法,而精度损失极小。但是实际代码实现中是直接在 FP8 下做 GEMM + ReLU 缩放。

2、Flash-Attention 风格的 Kernel 融合 (fp8_index):
DeepSeek 写了专门的 CUDA Kernel。在这个 Kernel 里,计算内积 + 寻找 Top-K 是在 GPU 的 SRAM(片上内存)里流式完成的。

1、它不需要把巨大的 $L \times L$ 分数矩阵写回显存(HBM)。
2、算出一段分数,就立马更新 Top-K 的候选集,算完的同时,Top-K 的索引也就选出来了。

4.1、Hadamard 旋转与 FP8 量化

DSA在实际代码实现中是直接在 FP8 下做 GEMM + ReLU 缩放,但这里还是简单介绍一下Hadamard旋转。

Hadamard 旋转,在低精度计算领域通常指的是 Hadamard 变换,它是一种正交变换, 可以被理解成一种“信息搅拌机”,通常大语言模型(LLM)激活值有个特点:少数几个固定维度上的值会异常巨大,成为 “异常值”(Outliers)。当想用只有少数位数的 FP8 来存这些数时,由于FP8 的动态范围很有限,就会陷入两难:

1、如果为了包容那个巨大的异常值,把整个张量的缩放因子调得很大。那么绝大多数的、数值正常的元素就会被“饿死”,被粗暴地舍入到0,大量信息丢失。
2、如果为了保护大多数正常值,把缩放因子调小。那个异常值又会“撑死”,超出最大表示范围变成NaN(非数)。

这就像要把一条鲸鱼和几百万条沙丁鱼装进同一个规定尺寸的鱼缸,无论怎么调整鱼缸的“缩放尺”,都很难兼顾。

Hadamard 旋转能把能量抹匀,Hadamard 变换矩阵 $\bm{H}$是一个充满 +1 和 -1 的正交方阵。它对向量 $\bm{x}$ 的旋转操作就是乘法 $\bm{H}\bm{x}$ 。它的魔法在于让信息“去局部化”。向量 $\bm{x}$ 中的每个元素,哪怕原先是0,经过变换后,在 $\bm{H}\bm{x}$ 的每一个新维度上,都变成了原向量所有元素的加权和(权重只有+1或-1)。这带来了一个直接后果:原来集中在少数维度的巨大能量(异常值),被均匀地“涂抹”到了旋转后的所有维度上。所有维度的数值范围变得非常接近,分布极其均匀,几乎不再有突出的异常值。一旦向量的各个维度分布均匀,之前的问题就迎刃而解。

这样可能要问:把输入 $\bm{x}$旋转了,那计算结果不就不对了吗?这就要说到一个巧妙的数学恒等式。由于 Hadamard 矩阵是对称且正交的, $\bm{H}^{-1}=\bm{H}$, 对于神经网络常见的线性层 $\bm{Y} = \bm{X}\bm{W}$:可以重写为

$$\begin{align}
\bm{Y} = \bm{X}\bm{H}\bm{H}^{-1}\bm{W}
\end{align}$$

这意味着:

1、离线操作:我们可以提前把权重矩阵也乘以 Hadamard 矩阵,得到新的权重 $\bm{W}:=\bm{H}^{-1}\bm{W}$。这步不占用推理时间。
2、在线操作:推理时,我们只对输入实时做 Hadamard 旋转,得到分布均匀的 $\bm{X}\bm{H}$。
3、快速计算:新的 $\bm{X}$ 和 $\bm{W}$ 的分布都很均匀,我们就可以安全、高效地用 FP8 计算 $\bm{Y}$,得到与原始高精度计算等价的结果。

更重要的是,得益于快速沃尔什-哈达玛变换(Fast Walsh-Hadamard Transform, FWHT),$\bm{X}\bm{H}$ 的计算复杂度仅为 $O(n \log n)$,远低于后续矩阵乘法的 $O(n^3)$,几乎可以忽略不计。

Hadamard 旋转能在几乎零额外开销下“熨平”异常值,离不开 快速沃尔什‑哈达玛变换 (Fast Walsh‑Hadamard Transform, FWHT)。下面我们从 Hadamard 矩阵的数学构造出发,推导旋转公式,并重点解释 FWHT 的蝶形算法如何实现 $O(n \log n)$ 复杂度,以及为何这相比矩阵乘法真的可以“忽略不计”。

4.2、稀疏索引的FP8融合算子

fp8_index的实现 借用了 FlashAttention 的 Tiling(分块) 和 Fusion(融合) 思想,解决了计算 $L\times L$ 相似度矩阵太慢、显存爆炸的问题,同时利用 FP8 量化 和 Hadamard 变换 解决了精度问题。

4.2.1、 核心痛点:为什么不能直接算?

如果按照标准 PyTorch 代码写 DSA 的 Indexer:

1
2
3
# 伪代码:标准实现
S = torch.matmul(Q_idx, K_idx.transpose(-2, -1)) # [Batch, L, L] 巨大!
mask = torch.topk(S, k=2048) # 在巨大的矩阵上找 Top-K

这在长上下文场景下(比如 $L=128K$)是不可行的:

1.显存爆炸:存储 $S$ 需要 $128K \times 128K \times 2$ Bytes $\approx 32GB$ 显存(仅一个 Head)。
2.计算浪费:算出所有 $L^2$ 个分数,最后却只留 Top-K 个,绝大部分计算被浪费。
3.Top-K 慢:在 CPU 或 GPU 上对 128K 长度的数组做全排序或 Top-K 也很耗时。

fp8_index Kernel 的目标:不生成巨大的 $S$ 矩阵,而是在 GPU 的 SRAM(片上内存) 里流式计算分数,算出一个丢一个,只把满足 Top-K 条件的留下来。

4.2.2、Lightning Indexer

在 DeepSeek-V3 的 DSA(DeepSeek Sparse Attention)[^2]机制中,Lightning Indexer 的核心逻辑还包含一个可学习的权重加权求和过程。

$$\begin{align}
S[t,s]=\sum_{h=1}^{H} w[t,h] \cdot \mathrm{ReLU}\big(\bm{q}^\T[t,h]\cdot\bm{k}[s]\big)
\end{align}$$

公式中的三个下标 $t, s, h$ 分别代表:

1、$t$ (Target / Query 索引),代表谁在看。它是序列中的某个位置(比如第 1 个词)。变化范围:$1$ 到 $L$。其中 $\dim[\bm{q}[t,h]]=d\times 1$
2、$s$ (Source / Key 索引),代表被看的是谁。它也是序列中的某个位置(比如第 3 个词)。变化范围:$1$ 到 $L$(在因果掩码下通常是 $s\le t$ )。其中 $\dim[\bm{k}[s]]=d\times 1$
3、$h$ (Head 索引)代表“从哪个角度看”。DSA 有多个 Indexer 头,每个头关注不同的特征(有的看局部,有的看语义)。变化范围:$1$ 到 $H$。

1、多头加权:$H$ 是 Indexer 的头数。不同于标准 Attention 直接将所有头的结果拼接,Indexer 会对所有头的分数进行加权求和,将多个相关性信号压缩为一个单一的“索引分数”。
2、ReLU 激活:引入 $\mathrm{ReLU}$ 是为了引入稀疏性和非线性,使得 Indexer 能够学习到更稀疏的选择模式(即只关注正相关的相关性)。
3、权重来源:$w_{t,h}$ 通常是一个小的可学习向量(或由 Query 状态投影得到),它让模型学会“在这个 Query 位置,哪个 Indexer 头的意见更重要”。

所以核心数学表达是将复杂的注意力筛选问题,简化为一个清晰、可微、可高效并行计算的表达式:

$$\begin{align}
\text{最终分数}=\sum_{\text{各个视角}}(\text{动态权重}\times{\text{只保留正相关的评分})
\end{align}$$

4.2.3、核心算法

再熟悉一下符号

  1. $\bm{Q}_{idx} \in \mathbb{R}^{L \times H \times d}$,其中$\bm{q}[t,h] \in \mathbb{R}^{d \times 1}$ 表示第 $t$ 个 token 在第 $h$ 个头的 Query 向量。
  2. $\bm{K}_{idx} \in \mathbb{R}^{L \times d}$ (假设 Key 在头维度共享,这是 DeepSeek MLA 的常见设计, 其中 $\bm{k}_s \in \mathbb{R}^{d \times 1}$ 表示第 $s$ 个 token 的 Key 向量。
  3. $\bm{W} \in \mathbb{R}^{L \times H}$, 其中$w_{t,j}$ 表示第 $t$ 个 token 对第 $j$ 个头的权重。

使用爱因斯坦求和约定一次性计算:
$$
\bm{S} =\text{einsum}_{th,ths \to ts}\bigg(\bm{W},\text{ReLU}\big(\text{einsum}_{thd,ds \to ths}(\mathbf{Q},\mathbf{K}^\T)\big)\bigg)
$$

Lightning Indexer 的矩阵形式可以写作:
$$\begin{align}
\mathbf{S} = \sum_{h=1}^{H} \left( \mathbf{W}[:,h] \odot \text{ReLU}\left( \bm{Q}[h]\cdot\bm{K}^\T \right) \right)
\end{align}$$

PyTorch 风格的伪代码实现

1
2
3
4
5
6
7
8
9
10
11
12
13
# Q_idx: [L, H_I, d_idx]
# K_idx: [L, d_idx]
# W : [L, H_I]
# 1. 计算相关性分数: [L, H_I, L]
# 注意:这里利用了广播机制
scores = torch.einsum('tjd,sd->tjs', Q_idx, K_idx)
# 2. 应用 ReLU
relu_scores = torch.relu(scores)
# 3. 加权求和: [L, H_I, L] * [L, H_I, 1] -> [L, L]
# 将 W 扩展为 [L, H_I, 1] 以便广播乘法,然后在 H_I 维度求和
I_matrix = torch.einsum('tj,tjs->ts', W, relu_scores)
# 或者分步写法:
# I_matrix = (W.unsqueeze(-1) * relu_scores).sum(dim=1)

4.2.4、分块计算

如果按标准 PyTorch 写在长上下文(例如 $L = 128K$)下,这根本跑不起来:

  • 显存:要完整存一个 $[L, H, L]$ 的中间张量(即使 H 很小,也远超 HBM);
  • 算力:算出 $L^2$ 个分数,最后却只用 Top-K,绝大部分计算是浪费。

生产级实现(DeepSeek 参考实现、NVIDIA cuDNN、vLLM/DeepGEMM)的核心思路是:

在 GPU 的 SRAM 里,按 tile 流式计算分数,边算边维护 Top-K 堆,绝不把整个 $\bm{S}$ 写到 HBM。

目标:对每个 Query 位置 $t$,在因果约束 $[0, \text{seqlen}_t)$ 内选出 Top-K 索引 $\mathcal{I}_t$,而不是显式构造完整的 $I[t,s]$。

步骤 0:量化 $\displaystyle \bm{Q}: = \mathrm{Quant}_{\text{FP8}}(\bm{Q}), \quad
\bm{K}: = \mathrm{Quant}_{\text{FP8}}(\bm{K})$

步骤 1:在 SRAM 中分块计算分数矩阵
把序列维度 $L$ 划分为若干块:Query 方向块大小:$B_q$;Key 方向块大小:$B_k$(通常 $B_k \ll L$,例如 64–256)。对每个 Query 块 $B_q \subseteq \{1,\dots,L\}$ 和每个 Key 块 $B_k \subseteq \{1,\dots,L\}$。

1、从 HBM 加载到 SRAM: $\displaystyle \bm{Q}_{B_q} \in \mathbb{R}^{B_q \times H \times d},\quad
\bm{K}_{B_k} \in \mathbb{R}^{B_k \times d}$

2、在 SRAM 中计算这一块的内积张量: $\displaystyle \bm{R}_{B_q\times B_k}[h] = \bm{Q}_{B_q}[h] \bm{K}_{B_k}^\T$ 。即对每个头 $h$ 做一次 $[B_q \times d] \times [d \times B_k] \rightarrow [B_q \times B_k]$ 的 FP8 GEMM。
3、对每个头应用 ReLU 和权重加权: $\displaystyle \bm{S}_{B_q\times B_k}
= \sum_{h=1}^{H} \bm{W}[:,h] \odot \mathrm{ReLU}!\big(\bm{R}_{B_q\times B_k}[h]\big)$ 。这里 $\odot$ 是逐元素乘法并沿头维度求和,得到 $[B_q \times B_k]$ 的分数块。

关键点:

  • $\bm{S}[B_q,B_k]$ 只在 SRAM 中存在,从不写回 HBM;
  • 每个 $t \in B_q$ 对 $B_k$ 内的 $s$ 都有一行分数。

步骤 2:在 SRAM 中维护 Top-K 堆并更新

1、对每个 Query 位置 $t \in B_q$,我们在 SRAM 中维护一个 最小堆 $\text{Heap}_t$:堆中保存 $(\text{score}, s)$ 对,堆容量为 $K$(例如 2048, 初始为空)。 对当前 Key 块 $B_k$,我们拿到 $\bm{S}[B_q,B_k]$ 的一行: $\displaystyle \{S[t,s]\}_{s \in B_k}$

2、对于每个满足因果约束的 $s$(例如 $s \le t$),执行堆更新:

$$\begin{align}
\text{Heap}_t \leftarrow \mathop{UpdateTopK}_{s}(\text{Heap}_t,\, S[t,s],\, s)
\end{align}$$

其中直观理解:若堆未满:直接插入 $(S[t,s], s)$;若堆已满:若 $S[t,s] > {\small\text{堆顶分数}}$,则替换堆顶,否则丢弃。这样,分数一旦算出就立刻用于更新堆,之后不需要再存储这个分数

步骤 3:写出 Top-K 索引

1、当所有 Key 块 遍历完毕后,每个 $t$ 的堆 $\text{Heap}_t$ 中已经保存了 Top-K 索引集合 $\mathcal{I}_t$: $\displaystyle \mathcal{I}_t = \mathrm{GetIndices}(\text{Heap}_t)$
2、将这些索引写回 HBM,作为后续主注意力(MLA)的稀疏掩码: $\displaystyle \bm{M}_{t,s} =
\begin{cases}
1, & s \in \mathcal{I}_t,\\
0, & {\small\text{其他}}
\end{cases}$

4.2.5、显存占用时间线

1、初始化分配堆空间 $\displaystyle \to B_q\times K\times 6$, 其中索引 = INT32 (4 Bytes),分数 = FP16 (2 Bytes) (平衡精度与显存)
2、加载FP8 Q/K 块 $\displaystyle \to B_q\times H \times d+ B_k\times d + B_q\times K\times 6$
3、计算相关分数(QK内积) $\mathop{ReLU}(\bm{R})$ $\displaystyle \to B_q\times B_k \times H + B_q\times H \times d+ B_k\times d + B_q\times K\times 6$
4、加载权重FP16 $\bm{W}$块 $\displaystyle \to B_q\times H \times 2 + B_q\times B_k \times H + B_q\times K\times 6$
4、FP16分配分数累加器 $\displaystyle \to B_q\times B_k \times 2 + B_q\times H \times 2 + B_q\times B_k \times H + B_q\times K\times 6$
5、更新 Top-K 堆 $\displaystyle \to B_q\times B_k \times 2 + B_q\times K\times 6$
6、计算结束 $\displaystyle \to B_q\times K\times 6$

通过,SRAM 的峰值,可以计算合理的 $B_q$ 和 $B_k$。

4.2.5、伪代码:概念级“fp8_index”Kernel
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
def fp8_index_conceptual(Q_idx, K_idx, W, K_topk, causal_ks, causal_ke):
"""
Q_idx: [L, H, d] (FP8 或已量化的低秩表示)
K_idx: [L, d] (FP8)
W : [L, H] (权重)
K_topk: 整数,例如 2048
causal_ks, causal_ke: [L] 的因果区间 [ks[t], ke[t])
"""
L = Q_idx.shape[0]
device = Q_idx.device
# 输出:每个 query 的 top-K 索引
topk_inds = torch.empty(L, K_topk, dtype=torch.int32, device=device)
# 对每个 query 独立处理(真实 kernel 会 batch + tile 并行)
for t in range(L):
# 当前 query 的因果有效 KV 区间
start = causal_ks[t]
end = causal_ke[t]
# 在 SRAM 中维护一个最小堆(按 score 排序,保留最大的 K_topk)
# 这里用 Python 堆示意,真实 kernel 在寄存器/共享内存中实现
heap = []
# 遍历 KV 区间,按 tile 划分
T_k = 64 # 示例 tile 大小
for s_start in range(start, end, T_k):
s_end = min(s_start + T_k, end)
# 1. 从 HBM 加载 Q[t] 和 K[s_start:s_end] 到 SRAM
q_t = Q_idx[t] # [H, d]
k_tile = K_idx[s_start:s_end] # [T_k, d]
# 2. 对每个头计算内积并应用 ReLU + 权重求和
# 这里用矩阵乘法示意,真实 kernel 会用 FP8 GEMM
# scores_tile: [T_k]
scores_tile = torch.einsum('hd,sd->s', q_t, k_tile) # 每个头的内积
scores_tile = torch.relu(scores_tile) # ReLU
# 加权求和:假设 W[t] 在 SRAM 中
I_ts = (W[t] * scores_tile).sum() # 标量分数
# 3. 对 s in s_start:s_end,更新堆
for offset, s in enumerate(range(s_start, s_end)):
score = I_ts[offset] # 当前 (t,s) 的分数
# 更新 Top-K 堆
if len(heap) < K_topk:
heappush(heap, (score, s))
else:
if score > heap[0][0]:
heapreplace(heap, (score, s))
# 4. 将堆中的索引写入输出
# 堆中元素从最小到最大,需要反转
topk_inds[t] = sorted([idx for (score, idx) in heap], reverse=True)
return topk_inds

说明

上层 PyTorch 代码看不到“堆”和“SRAM”,但在 CUDA/Triton/TileLang kernel 层,就是在寄存器/共享内存里维护一个 Top-K 结构;
cuDNN 的 DSA 模块把整个过程拆成:
IndexerForward:在 SRAM 中分块计算分数,并不写出完整 $S$;
IndexerTopK:基于分数做 radix Top-K,per-row 有效长度由 seq_lens 控制。
vLLM/DeepGEMM 目前先物化 logits(分数)张量,再用 fused Top-K kernel,但他们已经计划参考 DeepSeek 的 TileLang kernel 做进一步融合。

参考文献
[^1]: DeepSeek-AI. (2025, September 29). Boosting long-context efficiency with DeepSeek sparse attention [Technical report]. GitHub. https://github.com/deepseek-ai/DeepSeek-V3.2-Exp/blob/main/DeepSeek_V3_2.pdf
[^2]: Cai, Z., et al. (2025, December 2). DeepSeek-V3.2: Pushing the frontier of open large language models. arXiv. https://doi.org/10.48550/arXiv.2512.02556


版权声明
引线小白创作并维护的柠檬CC博客采用署名-非商业-禁止演绎4.0国际许可证。
本文首发于柠檬CC [ https://www.limoncc.com ] , 版权所有、侵权必究。
本文永久链接https://www.limoncc.com/post/1540fb65f18f359a/
如果您需要引用本文,请参考:
引线小白. (May. 25, 2026). 《大语言模型研究16——注意力机制优化之稀疏注意》[Blog post]. Retrieved from https://www.limoncc.com/post/1540fb65f18f359a
@online{limoncc-1540fb65f18f359a,
title={大语言模型研究16——注意力机制优化之稀疏注意},
author={引线小白},
year={2026},
month={May},
date={25},
url={\url{https://www.limoncc.com/post/1540fb65f18f359a}},
}

'