Contents

MiniPaged-Qwen:面向 Qwen3 推理的 Prefill FlashAttention

在大模型推理中,用户输入一段 Prompt 后,模型并不是立即进入逐 Token 生成。完整推理通常可以分成两个阶段:PrefillDecode

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 $$

/images/minipaged_qwen_prefill/01_attention_hbm_bottleneck.png

问题在于:$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} $$

/images/minipaged_qwen_prefill/02_flashattention_v2_traversal.png

这样就不需要一次构造完整的 $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 在此基础上进一步关注并行任务如何分配。它主要从三个方向改善:

  1. 减少非矩阵乘法的 FLOPs;
  2. 除了 Batch 和 Head 维度,还沿 Sequence Length 方向增加 Thread Block 级并行;
  3. 在单个 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。

/images/minipaged_qwen_prefill/03_qwen_grid_block_warp_mapping.png

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

/images/minipaged_qwen_prefill/04_raw_smem_layout.png

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} $$

一次迭代的数据流如下图。

/images/minipaged_qwen_prefill/05_kv_streaming_wmma_pipeline.png

下面逐步分析。


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_major

Row-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 Tilescores[16 * 16]
Online Maxrow_max[16]
Online Denominatorrow_sum[16]
Rescale Factoralpha[16]
Partial Probabilityprobs[16 * 16]
Partial Outputtile_out[16 * 128]
Running Output Stateoutput[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 的复刻。

当前实现仍然存在明显的进一步优化空间:

  • scoresprobstile_outoutput 使用了较大的 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 数据流的算法设计。


参考资料

  1. Tri Dao et al., FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness.
  2. Tri Dao, FlashAttention-2: Faster Attention with Better Parallelism and Work Partitioning.
  3. NVIDIA, CUDA Programming Guide,WMMA / Warp Matrix Functions 相关章节。
  4. Qwen Team, Qwen3-0.6B Model Configuration.
  5. MiniPaged-Qwen 项目中的 qwen_prefill_attention.cu、Qwen3 全层 Attention 集成测试与 Runtime 代码。