MiniPaged-Qwen:面向 Qwen3 推理的 Prefill FlashAttention
在大模型推理中,用户输入一段 Prompt 后,模型并不是立即进入逐 Token 生成。完整推理通常可以分成两个阶段:Prefill 和 Decode。
Prefill 一次处理整段 Prompt,拥有较长的 Query 序列,Attention 更接近矩阵乘法;Decode 每一步通常只有一个 Query Token,却需要不断读取增长的历史 KV Cache。两者虽然都在计算 Attention,但计算形态和优化重点完全不同。
MiniPaged-Qwen 项目的 Prefill 路径,目标是针对 Qwen3 的形状特征实现一个教学性较强的 FlashAttention-style CUDA Kernel:
- 输入采用 Qwen3 的 GQA 头布局;
- Q 按 Query 行分块;
- K/V 以 Tile 的方式流式遍历;
- 使用 Online Softmax 保证分块计算与完整 Softmax 数学等价;
- 使用 WMMA 将 $QK^T$ 和 $PV$ 映射到 Tensor Core;
- 所有中间 Score、Probability 和 Output Accumulator 保留在片上存储中,不在 HBM 中显式落盘完整的 $N\times N$ Attention Matrix。
本文先从 Transformer 和自注意力机制出发,分析传统 Attention 的瓶颈;随后介绍 FlashAttention 和 FlashAttention v2 的核心思想并推导完整公式;最后逐层拆解 MiniPaged-Qwen 中 Qwen3 Prefill 快速路径的 Grid、Block、Warp 映射、raw_smem 内存组织、KV 遍历以及 WMMA 计算过程。
1. 从大模型到 Transformer
大语言模型的训练目标可以简单理解为:给定前面的 Token,预测下一个 Token。
假设一段 Token 序列为:
$$ x_1,x_2,\dots,x_N $$
自回归语言模型希望学习:
$$ p(x_{t+1}\mid x_1,x_2,\dots,x_t) $$
Transformer 的关键能力在于:序列中的每个位置都可以根据其他位置的信息动态计算自己的表示。在 Decoder-only LLM 中,由于使用 Causal Mask,第 $i$ 个 Token 只能观察:
$$ x_1,x_2,\dots,x_i $$
而不能看到未来 Token。
Transformer 中最核心的模块就是 Self-Attention。对于输入:
$$ X\in\mathbb{R}^{N\times d_{model}} $$
通过三个线性投影得到:
$$ Q=XW_Q $$
$$ K=XW_K $$
$$ V=XW_V $$
单个 Attention Head 的输出为:
$$ O=\operatorname{Softmax}\left(\frac{QK^T}{\sqrt d}+M\right)V $$
其中 $M$ 是 Mask Matrix。对于 Causal Attention:
$$ M_{ij}= \begin{cases} 0,&j\leq i\ -\infty,&j>i \end{cases} $$
整个计算可以拆成三步:
$$ S=\frac{QK^T}{\sqrt d}+M $$
$$ P=\operatorname{Softmax}(S) $$
$$ O=PV $$
从公式来看,这只是两个矩阵乘法中间夹一个 Softmax。但真正放到 GPU 上,问题并没有这么简单。
2. 传统 Attention 的瓶颈在哪里?
先看最直接的实现方式。
假设:
$$ Q,K,V\in\mathbb{R}^{N\times d} $$
首先计算:
$$ S=QK^T\in\mathbb{R}^{N\times N} $$
然后把 $S$ 写入显存,再执行 Softmax:
$$ P=\operatorname{Softmax}(S) $$
再把 $P$ 写入显存,最后计算:
$$ O=PV $$

问题在于:$Q,K,V,O$ 的规模是 $O(Nd)$,但 $S$ 和 $P$ 的规模都是:
$$ O(N^2) $$
当序列长度增长时,中间矩阵会迅速膨胀。
更重要的是,GPU 的计算单元和显存之间存在明显的存储层次:
HBM / Global Memory
↓
L2 Cache
↓
Shared Memory / L1
↓
Register
↓
CUDA Core / Tensor Core传统实现的问题不是矩阵乘法公式错误,而是计算流程要求中间结果反复穿过这条存储层次:
Q,K 读入
↓
计算 S
↓
S 写 HBM
↓
S 从 HBM 读出
↓
Softmax
↓
P 写 HBM
↓
P 从 HBM 读出
↓
与 V 相乘于是,一个非常自然的问题出现了:
我们真正需要的结果只是 $O$,是否一定要把完整的 $S$ 和 $P$ 存到 HBM?
答案是否定的。
3. FlashAttention:改变计算顺序,而不是改变数学结果
FlashAttention 的核心思想可以概括为:
Tiling + Online Softmax + Kernel Fusion。
它并没有近似 Attention,也没有修改:
$$ O=\operatorname{Softmax}\left(\frac{QK^T}{\sqrt d}+M\right)V $$
而是重新安排计算顺序,使 Score Tile、Softmax State 和 Partial Output 尽量停留在片上存储中。
3.1 基本分块
把 Q 沿序列方向划分为 Query Block:
$$ Q= \begin{bmatrix} Q_0\ Q_1\ \vdots\ Q_{T_q-1} \end{bmatrix} $$
其中:
$$ Q_i\in\mathbb{R}^{B_r\times d} $$
把 K 和 V 同样沿 Token 方向划分:
$$ K_j,V_j\in\mathbb{R}^{B_c\times d} $$
对于固定的一个 Query Block $Q_i$,依次遍历所有 KV Block:
$$ (K_0,V_0),(K_1,V_1),\dots,(K_{T_k-1},V_{T_k-1}) $$
每次只计算局部 Score Tile:
$$ S_{ij}=\frac{Q_iK_j^T}{\sqrt d}+M_{ij} $$
其中:
$$ S_{ij}\in\mathbb{R}^{B_r\times B_c} $$

这样就不需要一次构造完整的 $N\times N$ Score Matrix。
但还有一个问题:Softmax 的每个元素依赖整行的最大值和分母,如何在只看到局部 Tile 时得到全局正确结果?
答案是 Online Softmax。
4. Online Softmax:FlashAttention 的数学基础
对于一行 Score:
$$ s=[s_1,s_2,\dots,s_N] $$
普通 Softmax 为:
$$ p_j=\frac{e^{s_j}}{\sum_{k=1}^{N}e^{s_k}} $$
为了避免指数溢出,通常减去全局最大值:
$$ m=\max_j s_j $$
于是:
$$ p_j=\frac{e^{s_j-m}}{\sum_{k=1}^{N}e^{s_k-m}} $$
如果 Score 被拆成多个 Tile,我们可以维护三个状态:
- 当前行最大值 $m$;
- 当前行 Softmax 分母 $l$;
- 当前行未归一化输出累积量 $\tilde O$。
初始化:
$$ m^{(0)}=-\infty $$
$$ l^{(0)}=0 $$
$$ \tilde O^{(0)}=0 $$
假设当前处理第 $t$ 个 KV Tile,得到:
$$ S_t=\frac{Q_iK_t^T}{\sqrt d}+M_t $$
先求当前 Tile 每行最大值:
$$ m_t=\operatorname{rowmax}(S_t) $$
与历史最大值合并:
$$ m^{new}=\max(m^{old},m_t) $$
由于最大值发生变化,旧的分母和旧的输出累积量都需要重新缩放。定义:
$$ \alpha=e^{m^{old}-m^{new}} $$
当前 Tile 的未归一化概率为:
$$ \tilde P_t=e^{S_t-m^{new}} $$
更新分母:
$$ l^{new}=\alpha l^{old}+\operatorname{rowsum}(\tilde P_t) $$
更新未归一化输出:
$$ \tilde O^{new}=\alpha\tilde O^{old}+\tilde P_tV_t $$
遍历全部 KV Tile 后:
$$ O=\frac{\tilde O}{l} $$
除法是逐行广播。
这个递推非常重要,因为它把原本必须观察完整 Score Row 才能计算的 Softmax,转换成可以 Tile by Tile 更新的状态机。
4.1 为什么这个更新是正确的?
假设历史 Score 集合的最大值是 $m^{old}$,加入新 Tile 后最大值变为 $m^{new}$。
历史分母原来是:
$$ l^{old}=\sum_{j\in old}e^{s_j-m^{old}} $$
换到新基准 $m^{new}$ 后:
$$ \sum_{j\in old}e^{s_j-m^{new}}
\sum_{j\in old}e^{s_j-m^{old}}e^{m^{old}-m^{new}} $$
所以:
$$ \sum_{j\in old}e^{s_j-m^{new}}=\alpha l^{old} $$
再加上新 Tile:
$$ l^{new}=\alpha l^{old}+\sum_{j\in tile}e^{s_j-m^{new}} $$
输出累积量同理。
因此,FlashAttention 不是“近似 Softmax”,而是利用指数函数的平移关系,将完整 Softmax 重写成等价的流式递推。
5. 从 FlashAttention 到 FlashAttention v2
第一版 FlashAttention 解决了最核心的问题:避免完整 $N\times N$ 中间矩阵的 HBM I/O。
FlashAttention v2 在此基础上进一步关注并行任务如何分配。它主要从三个方向改善:
- 减少非矩阵乘法的 FLOPs;
- 除了 Batch 和 Head 维度,还沿 Sequence Length 方向增加 Thread Block 级并行;
- 在单个 Thread Block 内重新分配 Warp 工作,减少 Warp 之间通过 Shared Memory 交换中间结果的成本。
这里需要强调:FlashAttention v2 并不是重新发明一种新的 Attention 数学公式,它的本质依然是:
Q Block 固定
↓
遍历 KV Block
↓
QK^T
↓
Online Softmax
↓
PV
↓
更新局部 Output State变化主要发生在:
- 哪些 Q Block 分配给哪些 CTA;
- 一个 CTA 内 Warp 如何分工;
- Q、K、V、Score、Output State 放在哪一级存储;
- 如何让 Tensor Core 计算和数据搬运更匹配。
MiniPaged-Qwen 的 Qwen3 Prefill 快速路径,正是围绕这条主线实现的。
6. Qwen3 Prefill 的输入形状
MiniPaged-Qwen 的 Prefill Attention 接口接收:
$$ Q\in\mathbb{R}^{B\times H_q\times N_q\times d} $$
$$ K,V\in\mathbb{R}^{B\times H_{kv}\times N_k\times d} $$
对于项目使用的 Qwen3-0.6B 配置:
$$ H_q=16 $$
$$ H_{kv}=8 $$
$$ d=128 $$
因此使用 GQA 映射:
$$ group_size=\frac{H_q}{H_{kv}}=2 $$
第 $qh$ 个 Query Head 对应:
$$ kvh=\left\lfloor\frac{qh}{group_size}\right\rfloor $$
例如:
qh 0,1 → kvh 0
qh 2,3 → kvh 1
...
qh 14,15 → kvh 7这意味着两个 Query Head 共享同一个 K/V Head,但 Query 本身仍然分别计算 Attention 输出。
7. Fast Path 的线程层次:一个 Block 处理 32 行 Q
快速路径的 Kernel 启动配置为:
$$ grid.x=\left\lceil\frac{N_q}{32}\right\rceil $$
$$ grid.y=H_q $$
$$ grid.z=B $$
每个 Block:
$$ 64\ threads=2\ warps $$
每个 Warp 负责:
$$ 16\ Query\ rows $$
所以每个 Block 总共处理:
$$ 2\times16=32 $$
行 Q。

Kernel 中:
const int warp = threadIdx.x / 32;
const int lane = threadIdx.x % 32;
const int b = blockIdx.z;
const int qh = blockIdx.y;
const int kvh = qh / (q_heads / kv_heads);
const int q_start =
blockIdx.x * 32 + warp * 16;如果:
$$ blockIdx.x=0 $$
那么:
Warp 0: q_start = 0 → Q row 0~15
Warp 1: q_start = 16 → Q row 16~31如果:
$$ blockIdx.x=1 $$
那么:
Warp 0: Q row 32~47
Warp 1: Q row 48~63因此,一个 Warp 的 Q Tile 固定为:
$$ Q_{warp}\in\mathbb{R}^{16\times128} $$
这就是为什么结构体中有:
half q[16 * 128];8. raw_smem:每个 Warp 独占一块工作区
快速路径定义:
struct __align__(16) QwenMmaWarpStorage {
half q[16 * 128];
half kv[16 * 128];
float scores[16 * 16];
half probs[16 * 16];
float tile_out[16 * 128];
float output[16 * 128];
float row_max[16];
float row_sum[16];
float alpha[16];
};Kernel 中:
extern __shared__ __align__(16) unsigned char raw_smem[];
auto* ws =
reinterpret_cast<QwenMmaWarpStorage*>(raw_smem)
+ warp;因此:
raw_smem
│
├── QwenMmaWarpStorage[0] → Warp 0
│
└── QwenMmaWarpStorage[1] → Warp 1
8.1 每个字段占多少空间?
以 FP16 为 $2$ Bytes、FP32 为 $4$ Bytes:
| 字段 | 形状 | 类型 | 大小 |
|---|---|---|---|
q | $16\times128$ | FP16 | $4096$ Bytes |
kv | $16\times128$ | FP16 | $4096$ Bytes |
scores | $16\times16$ | FP32 | $1024$ Bytes |
probs | $16\times16$ | FP16 | $512$ Bytes |
tile_out | $16\times128$ | FP32 | $8192$ Bytes |
output | $16\times128$ | FP32 | $8192$ Bytes |
row_max | $16$ | FP32 | $64$ Bytes |
row_sum | $16$ | FP32 | $64$ Bytes |
alpha | $16$ | FP32 | $64$ Bytes |
一个 Warp Storage 总大小为:
$$ 26304\ Bytes=25.6875\ KiB $$
一个 Block 有两个 Warp,所以动态 Shared Memory 为:
$$ 2\times26304=52608\ Bytes $$
即:
$$ 51.375\ KiB $$
这也是当前教学实现的一个明显权衡:把大量中间状态放进 Shared Memory,代码结构直观,但 Shared Memory 使用量较大,会影响可同时驻留的 Block 数量,实际优化时需要结合目标 GPU 的资源限制重新设计 Tile 大小和中间状态布局。
9. Q Tile 是怎样搬进 Shared Memory 的?
每个 Warp 只在进入 KV 循环之前加载一次 Q:
for (int i = lane; i < 16 * 128; i += 32) {
const int row = i / 128;
const int d = i % 128;
const int qi = q_start + row;
ws->q[i] = qi < q_len
? query[...]
: __float2half(0.0f);
ws->output[i] = 0.0f;
}一个 Warp 有 $32$ 个 lane,需要搬运:
$$ 16\times128=2048 $$
个 FP16 元素。
平均每个 lane 处理:
$$ \frac{2048}{32}=64 $$
个元素。
对应数据量为:
$$ 2048\times2=4096\ Bytes $$
也就是每个 Warp 一次搬运 $4\ KiB$ Q 数据。
这块 Q Tile 会在整个 KV 内循环中反复复用:
Global Q
↓ 只加载一次
ws->q [16,128]
↓
KV Tile 0: QK_0^T
KV Tile 1: QK_1^T
KV Tile 2: QK_2^T
...这种“固定 Q、流式遍历 KV”的结构正对应 FlashAttention v2 的核心遍历方式。
9.1 为什么 lane < 16 初始化 Online Softmax 状态?
代码:
if (lane < 16) {
ws->row_max[lane] = -INFINITY;
ws->row_sum[lane] = 0.0f;
}原因是一个 Warp 虽然有 $32$ 个线程,但当前 Warp 只负责 $16$ 个 Query Row。
每个 Query Row 只需要一个:
$$ m_i $$
和:
$$ l_i $$
所以:
lane 0 → row 0
lane 1 → row 1
...
lane 15 → row 15而 lane $16\sim31$ 在这个初始化指令中没有工作。
这并不意味着整个 Kernel 是“一线程处理一行 Q”。只有在行级 Online Softmax 状态更新阶段,采用了这种简单映射;Q 加载和 WMMA 计算仍然由整个 Warp 协同完成。
10. 内循环:每次遍历 16 个 KV Token
外层结构是:
for (int k_start = 0;
k_start < valid_kv_len;
k_start += 16) {
// 1. Load K Tile
// 2. QK^T
// 3. Causal Mask + Online Softmax
// 4. Load V Tile
// 5. P @ V
// 6. Update output state
}因此 KV Tile 大小为:
$$ B_c=16 $$
一个 KV Tile:
$$ K_t,V_t\in\mathbb{R}^{16\times128} $$
一次迭代的数据流如下图。

下面逐步分析。
11. 第一步:Global K → Shared Memory
代码:
for (int i = lane; i < 16 * 128; i += 32) {
const int row = i / 128;
const int d = i % 128;
const int ki = k_start + row;
ws->kv[i] = ki < valid_kv_len
? key[...]
: __float2half(0.0f);
}每个 Warp 搬运:
$$ 16\times128 $$
个 FP16 K 元素,即:
$$ 4\ KiB $$
从:
Global Memory
↓
ws->kv Shared Memory注意 ws->kv 是一个复用缓冲区。当前阶段装 K,后面算完 Score 和 Softmax 后,同一个空间会被 V 覆盖。
12. 第二步:WMMA 计算 $QK^T$
当前:
$$ Q_{warp}\in\mathbb{R}^{16\times128} $$
$$ K_t\in\mathbb{R}^{16\times128} $$
需要计算:
$$ S_t=Q_{warp}K_t^T $$
结果:
$$ S_t\in\mathbb{R}^{16\times16} $$
代码:
wmma::fragment<
wmma::accumulator,
16, 16, 16,
float
> score_frag;
wmma::fill_fragment(score_frag, 0.0f);
#pragma unroll
for (int d = 0; d < 128; d += 16) {
wmma::fragment<
wmma::matrix_a,
16, 16, 16,
half,
wmma::row_major
> q_frag;
wmma::fragment<
wmma::matrix_b,
16, 16, 16,
half,
wmma::col_major
> k_frag;
wmma::load_matrix_sync(
q_frag,
ws->q + d,
128
);
wmma::load_matrix_sync(
k_frag,
ws->kv + d,
128
);
wmma::mma_sync(
score_frag,
q_frag,
k_frag,
score_frag
);
}12.1 fragment 是什么?
WMMA 的抽象计算为:
$$ D=AB+C $$
q_frag 表示矩阵 A 的一个 Warp-level Tile,k_frag 表示矩阵 B 的一个 Tile,score_frag 表示 FP32 Accumulator。
对于当前形状:
$$ A\in\mathbb{R}^{16\times16} $$
$$ B\in\mathbb{R}^{16\times16} $$
$$ C,D\in\mathbb{R}^{16\times16} $$
需要注意:fragment 不是一个普通的二维数组。它表示由整个 Warp 协同维护的矩阵 Tile,内部元素分布在 Warp 各线程的寄存器中。程序不应该依赖某个矩阵元素具体落在哪个 lane 的寄存器里。
12.2 为什么需要 8 次 mma_sync?
一次 WMMA 的 K 维为:
$$ K_{mma}=16 $$
但 Qwen3 Fast Path 的 Head Dimension 为:
$$ d=128 $$
所以:
$$ \frac{128}{16}=8 $$
次累加。
把 Q 和 K 沿 Head Dimension 分成:
$$ Q=[Q_0,Q_1,\dots,Q_7] $$
$$ K=[K_0,K_1,\dots,K_7] $$
每个子块:
$$ Q_r,K_r\in\mathbb{R}^{16\times16} $$
于是:
$$ QK^T
\sum_{r=0}^{7}Q_rK_r^T $$
score_frag 初始为零,8 次 mma_sync 不断累加:
$$ S^{(0)}=0 $$
$$ S^{(r+1)}=Q_rK_r^T+S^{(r)} $$
最终得到:
$$ S^{(8)}=QK^T $$
12.3 为什么 K 不需要显式转置?
Shared Memory 中的 K 是 Row Major:
$$ K\in\mathbb{R}^{16\times128} $$
需要的是:
$$ K^T\in\mathbb{R}^{128\times16} $$
代码将 k_frag 声明成:
wmma::matrix_b
wmma::col_majorRow-Major 的 K 与 Col-Major 视角下的 $K^T$ 具有相同地址映射。
Row Major K 的地址:
$$ addr(K[i,j])=i\cdot128+j $$
令:
$$ B=K^T $$
Col Major B 的地址:
$$ addr(B[j,i])=j+i\cdot128 $$
两者相同:
$$ i\cdot128+j=j+i\cdot128 $$
因此只改变加载时的矩阵解释方式,就可以避免真正的数据转置。
13. 第三步:Causal Mask 与 Online Softmax
score_frag 计算完成后,先写回当前 Warp 的 Shared Memory:
wmma::store_matrix_sync(
ws->scores,
score_frag,
16,
wmma::mem_row_major
);得到:
$$ ws->scores\in\mathbb{R}^{16\times16} $$
然后前 16 个 lane 分别处理一行:
if (lane < 16) {
const int qi = q_start + lane;
const int causal_limit =
kv_len - q_len + qi + 1;
...
}普通 Prefill 中:
$$ q_len=kv_len $$
所以:
$$ causal_limit=qi+1 $$
意味着第 $qi$ 行只能访问:
$$ 0,1,\dots,qi $$
如果 Query 和 KV 长度不同,这个表达式采用右下角对齐的 Causal 语义。
对于当前 Tile,每一行先计算:
$$ tile_max_i=\max_j S_{ij} $$
随后:
$$ m_i^{new}=\max(m_i^{old},tile_max_i) $$
$$ \alpha_i= \begin{cases} e^{m_i^{old}-m_i^{new}},&m_i^{old}\ \text{有效}\ 0,&m_i^{old}=-\infty \end{cases} $$
每个局部概率:
$$ p_{ij}=e^{S_{ij}-m_i^{new}} $$
存入 FP16:
ws->probs[lane * 16 + col]然后更新:
$$ l_i^{new}=\alpha_i l_i^{old}+\sum_jp_{ij} $$
这与前面推导的 Online Softmax 递推完全对应。
14. 第四步:复用 ws->kv,加载 V Tile
Score 和 Softmax 计算结束后,当前 K Tile 已经没有继续保留的必要。
因此代码直接覆盖:
for (int i = lane; i < 16 * 128; i += 32) {
...
ws->kv[i] = value[...];
}于是数据流是:
阶段 A:
ws->kv = K Tile
↓
计算 QK^T
阶段 B:
ws->kv = V Tile
↓
计算 P V这是一个很典型的 Shared Memory Lifetime 优化:两个生命周期不重叠的数据共用同一块存储。
每次 KV Tile 中,该 Warp 从 Global Memory 读取:
$$ 4\ KiB\ K $$
加上:
$$ 4\ KiB\ V $$
总共:
$$ 8\ KiB $$
Q 则在进入循环之前加载一次并持续复用。
15. 第五步:WMMA 计算 $PV$
当前:
$$ \tilde P_t\in\mathbb{R}^{16\times16} $$
$$ V_t\in\mathbb{R}^{16\times128} $$
需要:
$$ \tilde P_tV_t\in\mathbb{R}^{16\times128} $$
代码沿输出 Head Dimension 每 16 列计算一次:
#pragma unroll
for (int d = 0; d < 128; d += 16) {
wmma::fragment<
wmma::matrix_a,
16, 16, 16,
half,
wmma::row_major
> p_frag;
wmma::fragment<
wmma::matrix_b,
16, 16, 16,
half,
wmma::row_major
> v_frag;
wmma::fragment<
wmma::accumulator,
16, 16, 16,
float
> out_frag;
wmma::load_matrix_sync(
p_frag,
ws->probs,
16
);
wmma::load_matrix_sync(
v_frag,
ws->kv + d,
128
);
wmma::fill_fragment(out_frag, 0.0f);
wmma::mma_sync(
out_frag,
p_frag,
v_frag,
out_frag
);
wmma::store_matrix_sync(
ws->tile_out + d,
out_frag,
128,
wmma::mem_row_major
);
}这里和 $QK^T$ 有一个区别。
在 $QK^T$ 中,8 次 MMA 是沿 K 维累加到同一个 score_frag:
$$ [16,128]\times[128,16] $$
而在 $PV$ 中,K 维本来就是 $16$,8 次循环是在生成输出的不同列区间:
$$ [16,16]\times[16,16] $$
分别得到:
output columns 0~15
output columns 16~31
...
output columns 112~127最终拼成:
$$ 16\times128 $$
的 tile_out。
16. 第六步:更新跨 KV Tile 的输出累积量
代码:
for (int i = lane; i < 16 * 128; i += 32) {
const int row = i / 128;
ws->output[i] =
ws->output[i] * ws->alpha[row]
+ ws->tile_out[i];
}对应:
$$ \tilde O_i^{new}
\alpha_i\tilde O_i^{old} + \tilde P_iV_i $$
注意此时还没有除以 Softmax 分母。
当所有 KV Tile 遍历结束后:
output[...] =
ws->output[i] / ws->row_sum[row];对应:
$$ O_i=\frac{\tilde O_i}{l_i} $$
到这里,完整的 Prefill Attention 输出得到。
17. 把源码与 FlashAttention v2 一一对应
可以把理论中的概念和当前实现对应起来:
| FlashAttention v2 概念 | MiniPaged-Qwen 实现 |
|---|---|
| Q Block | 一个 Warp 的 $16$ 行 Q |
| CTA 级 Q 分块 | 一个 Block 共处理 $32$ 行 Q |
| KV Block | 每次 $16$ 个 K/V Token |
| Score Tile | scores[16 * 16] |
| Online Max | row_max[16] |
| Online Denominator | row_sum[16] |
| Rescale Factor | alpha[16] |
| Partial Probability | probs[16 * 16] |
| Partial Output | tile_out[16 * 128] |
| Running Output State | output[16 * 128] |
| Tensor Core $QK^T$ | 8 次累加 mma_sync |
| Tensor Core $PV$ | 8 个输出子块 mma_sync |
从计算顺序来看:
固定 Q Tile
↓
Load K Tile
↓
8 × WMMA → QK^T
↓
Mask + Scale
↓
Online Softmax
↓
Load V Tile
↓
8 × WMMA → PV
↓
Update Output State
↓
Next KV Tile这与 FlashAttention v2 的前向计算主线是一致的。
但需要诚实说明:这个快速路径是一个面向学习和 Qwen3 特定形状的 FlashAttention-style 实现,并不是官方 FlashAttention v2 Kernel 的复刻。
当前实现仍然存在明显的进一步优化空间:
scores、probs、tile_out和output使用了较大的 Shared Memory;- K 和 V 的 Global Memory Load 没有使用异步拷贝流水;
- 没有 Double Buffer 隐藏 KV Load 延迟;
- Online Softmax 仍由前 16 个 lane 执行行级标量循环;
- Shared Memory 占用较高,可能限制 Occupancy;
- Tile 大小固定为 $16\times16$,没有针对不同序列长度和 GPU 架构做 Autotune。
因此项目的价值主要在于:把 FlashAttention 的数学递推、CUDA 线程映射、Shared Memory 生命周期和 Tensor Core MMA 连接成一条清晰、可运行、可验证的完整路径。
18. 从数据搬运角度重新看整个 Kernel
现在回到最开始的问题:FlashAttention 到底优化了什么?
对于一个 Warp:
Q
加载一次:
$$ 16\times128\times2=4096\ Bytes $$
之后跨所有 KV Tile 复用。
每个 K Tile
$$ 16\times128\times2=4096\ Bytes $$
每个 V Tile
$$ 16\times128\times2=4096\ Bytes $$
也就是每个 KV Tile 读取:
$$ 8\ KiB $$
但以下中间状态全部停留在片上存储:
$$ S_t\in\mathbb{R}^{16\times16} $$
$$ \tilde P_t\in\mathbb{R}^{16\times16} $$
$$ \tilde O\in\mathbb{R}^{16\times128} $$
不会产生完整的:
$$ N\times N $$
Score 和 Probability Matrix 的 HBM 写回。
这就是 FlashAttention 最核心的思想:
不是减少 Attention 的必要数学计算,而是减少昂贵的存储层次间数据移动。
19. Prefill 为什么适合这种实现?
Prefill 阶段的 Query 长度为:
$$ N_q=N_{prompt} $$
有大量 Query Row 可以并行。
例如 Prompt 长度为 $1024$,则:
$$ grid.x=\left\lceil\frac{1024}{32}\right\rceil=32 $$
再乘 Query Head 和 Batch:
$$ 32\times H_q\times B $$
可以形成大量 CTA 级任务。
而 Decode 阶段通常:
$$ N_q=1 $$
这时 Q 侧几乎没有序列并行度,主要问题变成不断读取增长的 KV Cache。因此 MiniPaged-Qwen 将 Prefill 与 Decode 分成两条不同路径:
Prefill
→ FlashAttention-style tiled kernel
→ 建立 Prompt KV Cache
Decode
→ 单 Token Query
→ Paged Attention 读取分页 KV Cache这也是整个 MiniPaged-Qwen 项目的核心架构出发点。
20. 正确性验证:不仅比较一个随机矩阵
一个 CUDA Attention Kernel 能跑通并不代表它能正确接入真实模型。
项目中的 Prefill 集成测试做了两条路径:
Reference:
Qwen3 所有层
→ Hugging Face eager attention
→ logits
Custom:
Qwen3 所有层
→ Mini Qwen Prefill Attention
→ logits最终比较:
- 全量 logits 最大绝对误差;
- 平均绝对误差;
- 相对 $L_2$ 误差;
- 最后一个 Token logits 的 Cosine Similarity;
- Top-1 Token 是否一致;
- 每一个 Attention Layer 是否确实调用自定义实现。
这样的验证比单独测试:
random Q,K,V
→ CUDA Kernel
→ torch reference更严格,因为真实 Qwen3 会引入 QK Norm、RoPE、GQA、Layer-by-Layer Error Accumulation 等因素。
21. 总结
本文从 Transformer 的 Self-Attention 出发,逐步分析了传统 Attention 的 $N\times N$ 中间矩阵为什么会导致严重的 HBM I/O,并推导了 FlashAttention 中 Online Softmax 的核心更新:
$$ m^{new}=\max(m^{old},m_t) $$
$$ \alpha=e^{m^{old}-m^{new}} $$
$$ l^{new}=\alpha l^{old}+\operatorname{rowsum}(\tilde P_t) $$
$$ \tilde O^{new}=\alpha\tilde O^{old}+\tilde P_tV_t $$
随后结合 MiniPaged-Qwen 的 Qwen3 Prefill 快速路径,分析了完整 CUDA 执行过程:
Grid
↓
一个 Block:64 threads = 2 warps
↓
一个 Warp:16 行 Q
↓
Q Tile [16,128] 一次加载并复用
↓
循环遍历 KV Tile [16,128]
↓
8 × WMMA 累加完成 QK^T
↓
Causal Mask + Online Softmax
↓
复用 Shared Memory 加载 V
↓
8 个 WMMA 子块完成 PV
↓
更新 Online Output State
↓
遍历完成后统一除以 row_sum理解 FlashAttention 最重要的一点不是记住某个 CUDA API,而是理解下面这条完整链路:
$$ \boxed{ \text{Attention 数学公式} \rightarrow \text{Online Softmax 重排} \rightarrow \text{Tile 划分} \rightarrow \text{线程与 Warp 映射} \rightarrow \text{内存层次设计} \rightarrow \text{Tensor Core MMA} } $$
只有把这些层次连起来,才能真正理解为什么 FlashAttention 不只是一个“更快的 Softmax”,而是一种围绕 GPU Memory Hierarchy 重新组织整个 Attention 数据流的算法设计。
参考资料
- Tri Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.
- Tri Dao, FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning.
- NVIDIA, CUDA Programming Guide,WMMA / Warp Matrix Functions 相关章节。
- Qwen Team, Qwen3-0.6B Model Configuration.
- MiniPaged-Qwen 项目中的
qwen_prefill_attention.cu、Qwen3 全层 Attention 集成测试与 Runtime 代码。