跳转至

6. FlashAttention 分块逻辑

  1. 实现 online softmax 的一维版本:给定一个向量,仅遍历一次,计算数值稳定的 softmax,写出 Python 代码。

1.1 标准 softmax 的数值稳定性问题

给定向量 $ x = [x_1, x_2, \ldots, x_n] $,softmax 定义为:

$$ softmax(x_{i})=\frac{e^{x_{i}}}{\sum_{j=1}^{n}e^{x_{j}}} $$

由于指数函数增长极快,当 $ x_i $ 较大时(如 1000),直接计算 $ e^{x_i} $ 会导致上溢;当所有 $ x_i $ 都较小时,可能导致下溢。因此数值稳定的做法是先减去最大值 $ m = \max(x) $,再计算指数。标准实现需要三次遍历:

  1. 找最大值 $ m_{0} $

  2. 计算每个 $ e^{x_i-m} $ 并累加得到 $ l $。

  3. 计算每个 $ \frac{e^{z_{i}-m}}{l} $ 得到结果。

这种方案在 FlashAttention 中无法直接使用,因为分块计算时无法预先知道全局最大值。

1.2 Online softmax 算法

Online softmax 只需一次遍历,通过维护当前遇到的最大值 m 和指数和 l 来增量计算。其核心思想是:当遇到一个新的元素 x 时,如果它比当前最大值 m 大,则需要将之前累积的指数和 l 乘以 $ e^{m-x} $(因为之前的指数都是相对于旧的 m 计算的,现在需要缩放到新的 m),然后更新 m = x,再累加新的 $ e^{x-m} $。如果新元素不大于当前最大值,则直接累加 $ e^{x-m} $。最后统一归一化。

数学证明(正确性):

设已处理了前 $ k $ 个元素,当前最大值 $ m_k $,指数和 $ l_k = \sum_{i=1}^{k} e^{x_{i}-m_k} $。对于新元素 $ x_{k+1} $:

若 $ x_{k+1} > m_k $,则新的最大值 $ m_{k+1} = x_{k+1} $。旧的指数和需要重新缩放:

$$ l_{k+1}=l_{k}\cdot e^{m_{k}-m_{k+1}}+e^{x_{k+1}-m_{k+1}}=\sum_{i=1}^{k}e^{x_{i}-m_{k+1}}+1 $$


与直接使用 $ m_{k+1} $ 计算的结果一致。

若 $ x_{k+1} \leq m_k $,则 $ m_{k+1} = m_k $,直接累加:

$$ l_{k+1}=l_{k}+e^{x_{k+1}-m_{k}} $$

1.3 Python 实现

代码块

1 import math 2 3 def online_softmax(x): 4 "……" 5 一次遍历计算 softmax,数值稳定。 6 Args: x: list or 1D tensor 8 Returns: 9 list of softmax probabilities 10 "……" 11 m = -float('inf') # 当前最大值 12 l = 0.0 # 指数和(相对于当前 m) 13 out = [0.0] * len(x) 14 15 for i, val in enumerate(x): 16 if val > m: 17 # 发现新的最大值,需要缩放旧的 l 18 l += math.exp(m - val) # m 是旧的最大值,val 是新的 19 m = val 20 # 累加当前元素相对于当前 m 的指数 21 l += math.exp(val - m) 22 out[i] = math.exp(val - m) # 暂存未归一化值 23 24 # 归一化 25 for i in range(len(x)): 26 out[i] /= l 27 return out

向量化版本(利用 PyTorch,但无法避免多次遍历):

代码块

1 def online_softmax_torch(x: torch.Tensor) -> torch.Tensor: 2 m = torch.tensor(-float('inf')) 3 l = torch.tensor(0.0) 4 out = torch.empty_like(x) 5 for i, val in enumerate(x):


6 if val > m: 7 l += torch.exp(m - val) 8 m = val 9 l += torch.exp(val - m) 10 out[i] = torch.exp(val - m) 11 return out / l

在 FlashAttention 中,online softmax 应用于每一行的注意力分数,并且每个 Q 块需要维护 running max 和 running sum,当遍历 K/V 块时进行增量更新。

2. 用公式推导标准 Attention 与 FlashAttention 的 I/O 复杂度,解释为何分块能减少 HBM 读写次数。

2.1 标准 Attention 的 I/O 分析

假设单头注意力,序列长度 N,头维度 d,数据以 FP16 存储(2 bytes)。在 GPU 上,HBM 带宽远低于计算吞吐,因此优化 HBM 访问是关键。

标准计算步骤及 HBM 读写(忽略 batch size 和 heads,它们可被类似分析):

  1. 从 HBM 读取 $ Q \in \mathbb{R}^{N \times d} $ 和 $ K \in \mathbb{R}^{N \times d} $,计算 $ S = QK^T $:读取 2Nd 个元素,输出 $ S \in \mathbb{R}^{N \times N} $ 需写回 $ N^2 $ 个元素。

  2. Softmax: 读取 $ S(N^2) $,写出 $ P \in \mathbb{R}^{N \times N}(N^2) $。

  3. 加权求和:读取 $ P(N^2) $ 和 $ V \in \mathbb{R}^{N \times d} $ ( $ Nd $),写出 $ O \in \mathbb{R}^{N \times d} $ ( $ Nd $)。

总 HBM 访问量 $ \approx O(Nd + N^2) $。当 $ N \gg d $ 时, $ N^2 $ 项占主导,且 $ S $ 和 $ P $ 矩阵的显存占用巨大。

2.2 FlashAttention 的分块策略与 I/O 复杂度

FlashAttention 将 Q, K, V 沿序列维度切分为块。设 Q 块大小为 $ B_r $, $ K/V $ 块大小为 $ B_c $,且 $ B_r, B_c = O(\sqrt{M}) $,其中 M 为 SRAM 容量(以元素个数计)。分块后,外层循环遍历 Q 块(共 $ T_r = \lceil N/B_r \rceil $ 个),内层循环遍历 $ K/V $ 块(共 $ T_c = \lceil N/B_c \rceil $ 个)。对每个 (Q 块, K 块) 对,在 SRAM 中计算局部注意力,利用 online softmax 增量更新输出,不将中间 S, P 写回 HBM。

每个(Q块,K块)对需要的HBM读取:

  • 加载 $ Q_i \in \mathbb{R}^{B_r \times d} $ ( $ B_r d $)

  • 加载 $ K_j, V_j \in \mathbb{R}^{B_e \times d} $ ( $ 2B_c d $)

其他统计量 m, l 很小可忽略。

输出 O 始终驻留在 SRAM 中,仅在处理完所有 K/V 块后才写回 HBM(每个 Q 块写回一次 $ O_i $,大小为 $ B_r d $)。


因此总 HBM 读取量 $ \approx T_r \times T_c \times (B_r d + 2B_c d) + T_r \times B_r d $(写回 O)。由于 $ T_r = N/B_r, T_c = N/B_c $,代入得:

HBM reads $ \approx \frac{N^2}{B_r B_c} (B_r d + 2B_c d) + N d = N^2 d \left( \frac{1}{B_c} + \frac{2}{B_r} \right) + N d $

若选择 $ B_r \approx B_c \approx \sqrt{M} $(通常设置 $ B_r = B_c $ 且与 SRAM 大小匹配),则 HBM 访问量为 $ O(N^2 d / \sqrt{M}) $。相比于标准 Attention 的 $ O(N^2) $,当 $ d \ll \sqrt{M} $ 时,这是一个巨大的降低。

实际中,FlashAttention 还通过 kernel 融合进一步减少 HBM 读写。例如,在 Decode 阶段 N 很小,Q 长度=1,分块优势不明显,但通过 PagedAttention 等其他机制优化。

  1. 用 PyTorch 编写分块注意力前向函数,输入 Q/K/V 及块大小 Bc、Br,采用 online softmax,写出内、外层循环的伪代码或代码。

3.1 单头示例(无 mask,支持 batch 和 head 维度已融合)

代码块

1 import torch 2 import math 3 4 def flash_attention_forward(Q, K, V, Br, Bc): 5 "" 6 Q, K, V: [N, d] (假设已经合并 batch 和 head,变成 [B*H, N, d] 亦可) 7 Br: Q 块大小 8 Bc: K/V 块大小 9 返回:O [N, d] 10 "" 11 N, d = Q.shape 12 scale = 1.0 / math.sqrt(d) 13 14 O = torch.zeros_like(Q) 15 L = torch.zeros(N, 1, device=Q.device, dtype=Q.dtype) # 最终的指数和 16 m = torch.full((N, 1), -float('inf'), device=Q.device, dtype=Q.dtype) # running max 17 18 Tr = (N + Br - 1) // Br 19 Tc = (N + Bc - 1) // Bc 20 21 for i in range(Tr): 22 q_start = i * Br 23 q_end = min(N, q_start + Br) 24 Qi = Q[q_start:q_end] # [Br, d]


Oi = O[q_start:q_end] # 引用,直接修改 mi = m[q_start:q_end] # running max for this block li = L[q_start:q_end] # running sum for this block

for j in range(Tc): k_start = j * Bc k_end = min(N, k_start + Bc) Kj = K[k_start:k_end] # [Bc, d] Vj = V[k_start:k_end]

# 计算  $ S = Qi \times Kj^T \times $ scale, shape [Br, Bc]
Sij = torch.matmul(Qi, Kj.T) * scale

# 更新 running max: mi_new = max(mi, row_max(Sij))
row_max_S = Sij.max(dim=1, keepdim=True).values
mi_new = torch.max(mi, row_max_S)

# 缩放旧的 Oi 和 li
Oi = Oi * torch.exp(mi - mi_new)
li = li * torch.exp(mi - mi_new)

# 当前块的 softmax 分子
Pij = torch.exp(Sij - mi_new)  # [Br, Bc]

# 更新 li 和 Oi
li = li + Pij.sum(dim=1, keepdim=True)
Oi = Oi + torch.matmul(Pij, Vj)

mi = mi_new

该 Q 块处理完毕,最终归一化并写回

O[q_start:q_end] = Oi / li L[q_start:q_end] = li m[q_start:q_end] = mi

return O


3.2 融合因果 mask 的修改

在计算 $ \text{Sij} $ 后,根据 Q 块和 K 块的全局位置构造 mask。因果 mask 确保 query 位置 p 只能看到 key 位置 $ k \leq p $。

python 复制 下载

在计算 $ S_{ij} $ 之后,更新 mi 之前插入:

q_pos = torch.arange(q_start, q_end, device=Q.device).unsqueeze(1) # [Br, 1] k_pos = torch.arange(k_start, k_end, device=Q.device).unsqueeze(0) # [1, Bc] causal_mask = (q_pos >= k_pos) # [Br, Bc], True 表示可见

$ S_{ij} $ = $ S_{ij.masked_fill} $(~causal_mask, float('-inf'))

如果整个 K 块都在 Q 块之后(即 $ k_start \geq q_end $),则可以直接跳过该 K 块以节省计算。

3.3 扩展到批量多头

输入形状 [B, H, N, d] 时,可以将 BH 视为独立的“序列”,将张量 reshape 为 [(BH), N, d],然后调用上述单头函数。更好的方法是直接在循环外部增加 batch/head 的并行维度(例如在 Triton 中)。

  1. 在上题基础上增加因果掩码(causal mask),说明分块计算时如何处理 mask,并修改代码逻辑。

4.1 因果掩码的数学形式

自回归注意力中,位置 i 只能关注位置 $ j \leq i $。掩码矩阵 M 定义为:

$$ M_{ij}=\begin{cases}0&j\leq i\ -\infty&j>i\end{cases} $$

在分块计算中,Q 块对应行范围 $ [q_start, q_end) $,K 块对应列范围 $ [k_start, k_end) $。根据这两个区间的关系,可分为三种情况:

  1. 完全屏蔽: $ k_start \geq q_end $,即整个K块的所有列都在Q块最大行的右侧。此时该K块对Q块完全不可见,可以直接跳过(continue)。

  2. 完全可见: $ k_end \leq q_start $,即 K 块所有列都在 Q 块最小行的左侧,全部可见,无需 mask。

  3. 部分可见:区间重叠。需要构造一个 [Br, Bc] 的布尔掩码,其中 True 表示对应位置 $ i \geq j $(即 Q 行索引 $ \geq K $ 列索引)。

4.2 修改后的内循环代码

代码块


内循环中:

for j in range(Tc): k_start = j * Bc k_end = min(N, k_start + Bc)

完全屏蔽检查

if k_start >= q_end: continue # 此 K 块全部在 Q 块之后,跳过

Kj = K[k_start:k_end] Vj = V[k_start:k_end] Sij = torch.matmul(Qi, Kj.T) * scale

构造因果掩码(仅当不是完全可见时)

if k_end > q_start: # 有部分屏蔽 q_pos = torch.arange(q_start, q_end, device=Q.device).unsqueeze(1) k_pos = torch.arange(k_start, k_end, device=Q.device).unsqueeze(0) causal_mask = (q_pos >= k_pos) Sij = Sij.masked_fill(~causal_mask, float('-inf'))

后续 online softmax 不变...

这种条件判断避免了不必要的掩码构造,提升了效率。

5. 实现 FlashAttention 前向中对单个 Q 块遍历所有 K/V 块的循环体:维护 running max、running sum 并更新局部输出,写出该核心循环的详细代码。

5.1 功能描述

给定一个 Q 块 Qi (形状 [Br, d]) 以及完整的 K, V (形状 [N, d]),使用分块 online softmax 计算该 Q 块对应的注意力输出 Oi 和最终的指数和 li。这是 FlashAttention 内层循环的核心。

5.2 完整代码

def flash_forward_q_block(Qi, K, V, Bc, scale, causal=False, q_start=0): """ Qi: [Br, d] 单个 Q 块 K, V: [N, d] Bc: K/V 块大小 scale: 缩放因子 1/sqrt(d) causal: 是否使用因果掩码 q_start: Q 块起始行索引(用于因果掩码)


返回: Oi [Br, d], li [Br, 1]

Br, d = Qi.shape N = K.shape[0] Oi = torch.zeros(Br, d, device=Qi.device, dtype=Qi.dtype) li = torch.zeros(Br, 1, device=Qi.device, dtype=Qi.dtype) mi = torch.full((Br, 1), -float('inf'), device=Qi.device, dtype=Qi.dtype)

Tc = (N + Bc - 1) // Bc q_end = q_start + Br

for j in range(Tc): k_start = j * Bc k_end = min(N, k_start + Bc)

# 因果掩码: 完全屏蔽跳过
if causal and k_start >= q_end:
    continue

Kj = K[k_start:k_end]  # [Bc, d]
Vj = V[k_start:k_end]

# 计算 Sij
Sij = torch.matmul(Qi, Kj.T) * scale  # [Br, Bc]

# 应用因果掩码
if causal and k_end > q_start:
        q_pos = torch.arange(q_start, q_end, device=Qi.device).unsqueeze(1)
        k_pos = torch.arange(k_start, k_end, device=Qi.device).unsqueeze(0)
        mask = q_pos >= k_pos
        Sij = Sij.masked_fill(~mask, float('-inf'))

# Online softmax 更新
row_max = Sij.max(dim=1, keepdim=True).values
mi_new = torch.max(mi, row_max)

# 缩放旧的 Oi 和 li
Oi = Oi * torch.exp(mi - mi_new)
li = li * torch.exp(mi - mi_new)

# 当前块的权重
Pij = torch.exp(Sij - mi_new)
li = li + Pij.sum(dim=1, keepdim=True)
Oi = Oi + torch.matmul(Pij, Vj)
mi = mi_new

56 0i = 0i / li 57 return 0i, li

5.3 注意事项

  • 当 causal=True 时,完全屏蔽的块跳过可以节省计算。
  • 在 PyTorch 实现中,循环内的 torch.exp 和 torch.matmul 是高效的,但在 GPU kernel 中需考虑向量化。
  • 真实的 FlashAttention kernel 会将多个 Q 块打包以增加并行度。

  • 编写函数,给定 Q/K/V 和块大小,返回注意力输出 O 和 log-sum-exp L,同时保存反向所需的 m、I 等中间统计量。

在 FlashAttention 前向过程中,为了节省显存,只保存每行的最终最大值 m 和最终指数和 L(即归一化分母),以及原始 Q、K、V。反向传播时,通过重计算注意力权重来得到梯度。因此我们需要在 forward 中额外返回这些统计量。

def flash_attention_forward_save_stats(Q, K, V, Br, Bc, causal=False):

Q, K, V: [N, d]
返回: O [N, d], L [N, 1], m [N, 1]

N, d = Q.shape
scale = 1.0 / math.sqrt(d)
O = torch.zeros_like(Q)
L = torch.zeros(N, 1, device=Q.device, dtype=Q.dtype)
m = torch.full((N, 1), -float('inf'), device=Q.device, dtype=Q.dtype)

Tr = (N + Br - 1) // Br
Tc = (N + Bc - 1) // Bc
for i in range(Tr):
    q_start = i * Br
    q_end = min(N, q_start + Br)
    Qi = Q[q_start:q_end]
    Oi = O[q_start:q_end]
    mi = m[q_start:q_end]
    li = L[q_start:q_end]
for j in range(Tc):
    k_start = j * Bc

k_end = min(N, k_start + Bc)

因果mask处理省略...

Kj = K[k_start:k_end] Vj = V[k_start:k_end] Sij = torch.matmul(Qi, Kj.T) * scale

... apply mask ...

row_max = Sij.max(dim=1, keepdim=True).values mi_new = torch.max(mi, row_max) Oi = Oi * torch.exp(mi - mi_new) li = li * torch.exp(mi - mi_new) Pij = torch.exp(Sij - mi_new) li = li + Pij.sum(dim=1, keepdim=True) Oi = Oi + torch.matmul(Pij, Vj) mi = mi_new

O[q_start:q_end] = Oi / li L[q_start:q_end] = li m[q_start:q_end] = mi

return O, L, m


在反向时,我们可以根据保存的 m 和 L,结合重新计算出的 Sij 和 Pij,来推导 dQ,dK,dV。

  1. 解释 FlashAttention 反向为何需要 L、m 等中间值,并写出对一对 (Q 块, K/V 块) 重计算注意力权重及梯度累加的伪代码。

7.1 反向传播公式

设前向输出 $ O = \text{softmax}(QK^T/\sqrt{d})V $。损失函数对 $ O $ 的梯度为 $ dO $。我们需要求 $ dQ, dK, dV $。推导如下: 定义 $ S = QK^T/\sqrt{d} $, $ P = \text{softmax}(S) $,则 $ O = PV $。 首先求 $ dV = P^T dO $。 然后求 $ dP = dOV^T $。 接着求 $ \text{softmax} $ 的梯度: $ dS = P \odot (dP - \text{rowsum}(dP \odot P)) $。 最后 $ dQ = (dS/\sqrt{d})K $, $ dK = (dS^T/\sqrt{d})Q $。

在分块计算中,我们无法存储完整的 \(P\),但可以在反向时通过保存的每行的最终最大值 \(m\) 和最终指数和 \(L\)(注意:此处的 \(L\)\(l = \sum_j e^{S_{ij} - m_i}\),不是最终的 \(\log\text{-sum} - \exp\)),以及重新计算得到的 \(S\),来恢复 \(P\)

$$ P_{ij}=\exp(S_{ij}-m_{i})/L_{i} $$

其中 $ m_{i} $ 是第 i 行的全局最大值, $ L_{i} $ 是全局指数和。因为前向已经保存了这些,所以反向时只需按块重新计算 $ S_{ij} $,就可以得到对应的 $ P_{ij} $,进而累加梯度。

7.2 单对 (Q块, K/V块) 的重计算及梯度累加伪代码

代码块

已知:dO_i [Br, d] 上游梯度(对应 Q 块 i),Q_i,K_j,V_j,

2 # m_saved_i [Br, 1] 保存的该 Q 块的全局最大值, 3 # L_saved_i [Br, 1] 保存的该 Q 块的全局指数和。 4 # 对于每个 K/V 块 j: 5 Sij = Q_i @ K_j.T * scale 6 # 恢复 Pij 7 Pij = torch.exp(Sij - m_saved_i) / L_saved_i # [Br, Bc] 8 # 计算 dV 贡献:dP = dO_i @ V_j.T,然后累加到 dV_j 9 dV_j += Pij.T @ dO_i 10 # 计算 dS 11 dPij = dO_i @ V_j.T # [Br, Bc] 12 dSij = Pij * (dPij - (dPij * Pij).sum(dim=1, keepdim=True)) 13 # 累加 dQ_i 和 dK_j 14 dQ_i += dSij @ K_j * scale 15 dK_j += dSij.T @ Q_i * scale


完整的反向传播需要在外层循环遍历 Q 块,内层循环遍历 K/V 块,并累加梯度。这就是 FlashAttention 高效反向的核心。

8. 用 Triton 语言实现简易版 FlashAttention 前向 kernel,包含块索引计算、SRAM 上 online softmax 和写回全局内存。

Triton 允许我们用类似 Python 的语法编写高性能 GPU kernel。以下给出一个仅支持单头、无 mask 的简化版示例,重点展示 online softmax 和分块逻辑。

代码块

1 import triton 2 import triton.language as tl 3 4 @triton.jit 5 def _flash_attention_fwd_kernel( 6 Q_ptr, K_ptr, V_ptr, O_ptr, L_ptr, M_ptr, 7 N, d, Br, Bc, scale, 8 stride_q_n, stride_k_n, stride_v_n, stride_o_n, 9 BLOCK_Q: tl.constexpr, BLOCK_KV: tl.constexpr, D: tl.constexpr 10 ): 11 # 每个 program 处理一个 Q 块 12 pid = tl.program_id(0) 13 q_start = pid * BLOCK_Q 14 if q_start >= N: 15 return 16 # 加载 Q 块 [BLOCK_Q, D] 到 SRAM 17 offs_q = tl.arange(0, BLOCK_Q) 18 offs_d = tl.arange(0, D) 19 Q = tl.load(Q_ptr + (q_start + offs_q[:, None]) * stride_q_n + offs_d[None, :], 20 mask = (q_start + offs_q[:, None]) < N, other=0.0) 21 22 # 初始化 O, l, m 23 O = tl.zeros([BLOCK_Q, D], dtype=tl.float32) 24 l = tl.zeros([BLOCK_Q, 1], dtype=tl.float32) 25 m = tl.full([BLOCK_Q, 1], float('-inf'), dtype=tl.float32) 26 27 # 遍历 K/V 块 28 Tc = (N + BLOCK_KV - 1) // BLOCK_KV 29 for j in range(Tc): 30 k_start = j * BLOCK_KV 31 offs_kv = tl.arange(0, BLOCK_KV) 32 # 加载 K, V 块 33 K = tl.load(K_ptr + (k_start + offs_kv[:, None]) * stride_k_n + offs_d[None, :],


mask=(k_start + offs_kv[:, None]) < N, other=0.0) V = tl.load(V_ptr + (k_start + offs_kv[:, None]) * stride_v_n + offs_d[None, :], mask=(k_start + offs_kv[:, None]) < N, other=0.0)

计算 $ S = Q \times K^T \times scale $

S = tl.dot(Q, K.T) * scale # [BLOCK_Q, BLOCK_KV]

因果 mask 可在此通过 tl.where 等处理

online softmax 更新

m_new = tl.maximum(m, tl.max(S, axis=1, keepdim=True))

缩放旧的 O 和 l

O = O * tl.exp(m - m_new) l = l * tl.exp(m - m_new) P = tl.exp(S - m_new) l = l + tl.sum(P, axis=1, keepdim=True) O = O + tl.dot(P, V) m = m_new

最终归一化

O = O / l

写回全局内存

tl.store(O_ptr + (q_start + offs_q[:, None]) * stride_o_n + offs_d[None, :], O, mask=(q_start + offs_q[:, None]) < N) tl.store(L_ptr + (q_start + offs_q), l, mask=(q_start + offs_q) < N) tl.store(M_ptr + (q_start + offs_q), m, mask=(q_start + offs_q) < N)

调用时需设置 grid 大小等。该 kernel 仅示意,真实实现还需处理多 batch、head,以及反向。

9. 实现支持多头批量处理的分块注意力:输入 [B, H, N, d],要求块循环对 batch 和 head 维度并行,写出关键代码结构。

为了最大化 GPU 并行度,我们通常将 $ \mathbf{B}*\mathbf{H} $ 视为独立的“序列”,每个线程块处理一个 (batch, head) 的一个 Q 块。在 Triton 中,通过增加 program_id 的维度来实现。

9.1 代码结构(Triton)

代码块

1 @triton.jit 2 def batched_flash_attention_kernel( 3 Q, K, V, O, L, M, 4 B, H, N, d, Br, Bc, scale, 5 stride_q_b, stride_q_h, stride_q_n, 6 ...


7 BLOCK_Q: tl.constexpr, BLOCK_KV: tl.constexpr 8 ): 9 # 使用 2D grid: (BH, num_blocks_q) 10 pid_bh = tl.program_id(0) # 合并的 batchhead 索引 11 pid_q = tl.program_id(1) # Q 块索引 12 b = pid_bh // H 13 h = pid_bh % H 14 q_start = pid_q * BLOCK_Q 15 # 偏移量计算:定位到对应的 Q, K, V 16 offs_b = b * stride_q_b + h * stride_q_h 17 # ... 加载 Q 块等,后续逻辑与单头相同

在 PyTorch 中,我们可以简单地将输入 view 成 $ (B\times H, N, d) $,然后循环处理每个 $ (B\times H) $ 样本的 Q 块。但这种纯 Python 循环效率很低,仅用于理解。

9.2 PyTorch 实现(概念性)

代码块

1 def flash_attention_batch(Q, K, V, Br, Bc): 2 B, H, N, d = Q.shape 3 Q = Q.view(BH, N, d) 4 K = K.view(BH, N, d) 5 V = V.view(BH, N, d) 6 O = torch.empty_like(Q) 7 L = torch.empty(BH, N, 1) 8 # 循环每个“序列” 9 for b in range(B*H): 10 O[b], L[b] = flash_attention_forward(Q[b], K[b], V[b], Br, Bc) 11 return O.view(B, H, N, d)

实际 GPU kernel 会将 B*H 并行化。

  1. 简述 FlashAttention-2 的改进,并写出将 Q 作为内循环的前向伪代码,对比 v1 的外循环区别。

10.1 FlashAttention-2的主要改进

  1. 循环顺序交换:v1 外循环遍历 Q 块,内循环遍历 K/V 块;v2 变为外循环遍历 K/V 块,内循环遍历 Q 块。这样 K/V 块只需加载一次,就可以服务所有 Q 块,减少了 HBM 读取 K/V 的次数。

  2. 更好的并行化:v2 在序列长度维度上提供更多并行度。在 decode 阶段(Q 长度=1),v1 只有很少的 Q 块,GPU 利用率低;v2 将 K/V 分块作为外层,可以有更多线程块并行。


  1. 减少非矩阵乘操作:通过优化循环,减少了 softmax 统计量的更新次数和寄存器压力,提升效率。

  2. 更优的 warp 调度:改进共享内存访问模式,减少 bank conflict。

10.2 v2 外循环遍历 K/V 块的伪代码

代码块

def flash_attention_v2_forward(Q, K, V, Br, Bc): N, d = Q.shape scale = 1.0 / math.sqrt(d) O = torch.zeros(N, d) L = torch.zeros(N, 1) m = torch.full((N, 1), -float('inf'))

Tc = (N + Bc - 1) // Bc
Tr = (N + Br - 1) // Br

for j in range(Tc):
    k_start = j * Bc
    k_end = min(N, k_start + Bc)
    Kj = K[k_start:k_end]
    Vj = V[k_start:k_end]

    for i in range(Tr):
        q_start = i * Br
        q_end = min(N, q_start + Br)
        Qi = Q[q_start:q_end]

    # 注意这里 Oi, mi, li 需要从全局内存加载并写回,因为多个 K 块会交错更新同一个 Q 块。
    # 在 v2 中,对于每个 K 块,所有 Q 块都会更新,因此 O, L, m 需要常驻全局或通过原子操作同步。
    # 实际实现中,每个 Q 块会维护自己的 Oi, mi, li,当处理完所有 K 块后再写回。
    # 但交换循环后,Q 块在 K 块循环内部,因此需要额外的同步。
    # 最简单的实现是将 O, L, m 放在全局内存,每次加载更新后写回,但这样 I/O 增加。
    # 真正的 v2 通过将 Q 分成更大的块或使用 warp 级并行来减少这种开销。
    return O / L

10.3 v1 与 v2 的区别总结

v1: for Qi: for Kj。优点:每个Q块的O,m,l可以完全保留在SRAM中,直到处理完所有K/V块才写回,HBM写入少。缺点:K/V块被重复加载多次,当Q块较多时K/V读取量大。

v2: for Kj: for Qi。优点:K/V块只需加载一次,减少HBM读取。缺点:Q块的中间状态不能完全留在SRAM,需要频繁加载/写回,但通过更大的Q块和高效的缓存策略来缓解。


适用场景:v2 在长序列 prefill 和 decode 阶段均表现出更好性能,已成为主流实现。FlashAttention-2 的改进使得其更适应现代 GPU 架构,尤其是在大 batch 和长序列时。

11. 在解码阶段每次只生成一个 token,设计 FlashDecoding 的核心流程:将 K/V 分块后各块独立计算 softmax,再合并,写出实现步骤或伪代码。

11.1 问题背景

在自回归解码阶段,Query 长度仅为 1,但 Key/Value 序列长度可能非常长(如 128k)。如果直接使用标准 Flash Attention,由于 Q 只有 1 个 token,外层 Q 块循环只有一个块,内层 K/V 循环仍然很长,但每个 K/V 块与单个 Q token 计算的点积仅为向量-矩阵乘法,计算密度极低,大量线程块被浪费,GPU 利用率低下。Flash Decoding 正是为了解决这一问题而提出的。

11.2 核心思想

FlashDecoding 将长序列的 K 和 V 分成多个块,并并行地为每个块独立计算该块内部的局部 softmax(包含局部最大值和局部指数和),然后将各个块的局部结果通过一次额外的“归约”(reduction)步骤合并,得到最终的注意力输出。这样做可以将原本串行的 K/V 遍历并行化,大幅提升长序列解码的 GPU 利用率。

11.3 算法步骤(以单个 Query token 为例)

设 Query 向量 $ q \in \mathbb{R}^d $,Key 和 Value 为 $ K, V \in \mathbb{R}^{N \times d} $,其中 $ N $ 很大。将 $ K, V $ 沿序列维度切分为 $ T_c $ 个块,每块大小 $ B_c $。

第一阶段:分块独立计算

对于每个 K/V 块 j:

  1. 加载 $ K_j \in \mathbb{R}^{B_e \times d} $,计算局部注意力分数 $ s_j = qK_j^T / \sqrt{d} \in \mathbb{R}^{1 \times B_e} $。

  2. 计算该块的局部最大值 $ m_{j} = \max(s_{j}) $

  3. 计算该块的局部未归一化权重 $ p_j = \exp(s_j - m_j) $,局部指数和 $ l_j = \sum p_j $。

  4. 计算该块对输出的局部贡献 $ o_j = p_j V_j \in \mathbb{R}^d $。

上述所有块的计算可以完全并行,因为块间无依赖。每个块输出一个三元组 $ (m_j, l_j, o_j) $。


第二阶段:归约合并

在所有块计算完成后,我们需要将这些局部结果合并,得到全局的 softmax 输出。这一步是 FlashDecoding 的关键,它使用 online softmax 的归约思想:

  1. 初始化全局最大值 $ m_{global} = -\infty $,全局指数和 $ l_{global} = 0 $,全局输出 $ o_{global} = 0 $。

  2. 顺序遍历每个块(或使用树形归约),对于块 $ j $:

  3. 新的全局最大值 $ m_{new} = \max(m_{global}, m_j) $。
  4. 缩放旧累加值: $ l_{global} \leftarrow l_{global} \cdot e^{m_{global} - m_{new}} + l_j \cdot e^{m_j - m_{new}} $。
  5. 缩放旧输出: $ o_{global} \leftarrow o_{global} \cdot e^{m_{global} - m_{new}} + o_j \cdot e^{m_j - m_{new}} $。
  6. 更新 $ m_{global} = m_{new} $。

  7. 最终输出 $ o = o_{global}/l_{global} $。

由于归约是顺序的,计算量相对于分块内计算很小,但仍然可以进一步优化:使用类似树形归约的方式在多级中进行,或者使用原子操作在 GPU 上实现。

11.4 伪代码

代码块

1 // Stage 1: 并行计算每个 KV 块的局部统计量 2 parallel for j in 0..Tc-1: load Kj, Vj 3 s = q * Kj^T / sqrt(d) // [1, Bc] 4 m_j = max(s) 5 p_j = exp(s - m_j) // [1, Bc] 6 l_j = sum(p_j) 7 o_j = p_j * Vj // [1, d] 8 store (m_j, l_j, o_j) 9 10 11 // Stage 2: 归并 12 m_global = -inf, l_global = 0, o_global = 0 13 for j in 0..Tc-1: 14 m_new = max(m_global, m_j) 15 l_global = l_global * exp(m_global - m_new) + l_j * exp(m_j - m_new) 16 o_global = o_global * exp(m_global - m_new) + o_j * exp(m_j - m_new) 17 m_global = m_new 18 output = o_global / l_global


11.5 与 FlashAttention 的对比

  • FlashAttention 是逐块串行处理 K/V,但 Q 也是分块的,适合 prefill 阶段。
  • FlashDecoding 利用了解码时 Q 只有 1 个 token 的特殊性,将 K/V 块完全并行,然后归约。这样可以将单个 query 的计算分布在多个线程块上,极大提升了长序列时的 GPU 占用率。

12. 当 KV Cache 采用 PagedAttention 物理块管理时,写出融合分页块表的 FlashAttention 前向流程:从块表加载 KV 块,完成分块 softmax,给出伪代码。

12.1 背景

PagedAttention 将 KV Cache 划分为固定大小的物理块,并通过逻辑块表将序列的逻辑位置映射到物理块。在注意力计算时,我们需要根据块表来加载 K 和 V。将分页块表与 FlashAttention 融合可以避免先拷贝出一个连续 KV 缓冲区,直接在 kernel 中按块表索引访问物理块。

12.2 融合流程

以单个序列的 prefill 为例(decode 类似),已知该序列的块表 block_table(长度为逻辑块数),物理 KV 缓存存储为全局内存中的大张量 kv_cache,形状为

[2, num_blocks, num_heads, block_size, head_dim]。假设我们处理一个头、一个Q块(索引为i),并循环所有逻辑KV块。

伪代码:

代码块

1 // grid: 每个线程块处理一个 Q 块 2 global void flash_paged_attention_kernel( 3 Q, block_table, kv_cache, O, L, m_global, 4 seq_len, block_size, head_dim, scale 5 ) { 6 int q_start = block_idx.x * Br; 7 int q_end = min(q_start + Br, seq_len); 8 // 加载 Q 块到共享内存 9 shared float Q_s[Br][head_dim]; 10 load Q[q_start:q_end] -> Q_s; 11 12 float O_s[Br][head_dim] = {0}; 13 float l_s[Br] = {0}; 14 float m_s[Br] = {-inf}; 15 16 int num_logic_blocks = ceil(seq_len / block_size); 17 for (int lb = 0; lb < num_logic_blocks; lb++) { 18 int phy_idx = block_table[lb]; // 从块表获取物理块索引


if (phy_idx < 0) continue; // 未分配则跳过

// 从全局物理缓存加载 K 和 V 块到共享内存 shared float K_s[block_size][head_dim]; shared float V_s[block_size][head_dim]; int k_start = lb * block_size; int k_end = min(k_start + block_size, seq_len); int valid_tokens = k_end - k_start; load kv_cache[0][phy_idx][head][0:valid_tokens] -> K_s; load kv_cache[1][phy_idx][head][0:valid_tokens] -> V_s; __syncthreads();

// 计算 S = Q_s * K_s^T * scale (只针对有效token) float S[Br][block_size]; for (int r = 0; r < q_end - q_start; r++) { for (int c = 0; c < valid_tokens; c++) { S[r][c] = dot(Q_s[r], K_s[c]) * scale; } } // 应用因果mask(如果 prefill 需要) // ...

// online softmax 更新 for (int r = 0; r < q_end - q_start; r++) { float row_max = max(S[r][0..valid_tokens]); float m_new = max(m_s[r], row_max); float factor = exp(m_s[r] - m_new); l_s[r] = l_s[r] * factor + sum(exp(S[r] - m_new)); // 更新 O_s[r] for (int d = 0; d < head_dim; d++) { O_s[r][d] = O_s[r][d] * factor; for (int c = 0; c < valid_tokens; c++) { O_s[r][d] += exp(S[r][c] - m_new) * V_s[c][d]; } } }

// 最终归一化并写回全局内存 for (int r = 0; r < q_end - q_start; r++) { for (int d = 0; d < head_dim; d++) O[q_start + r][d] = O_s[r][d] / l_s[r]; }


说明:此 kernel 中,K 和 V 不是连续的逻辑序列,而是通过块表 $ \underline{\text{block_table[lb]}} $ 映射到物理块索引,再根据物理索引从 $ \underline{\text{kv_cache}} $ 中加载。这样就实现了 PagedAttention 与 FlashAttention 的无缝融合,无需额外拷贝。

13. 实现数值稳定的分块 softmax 内核,演示若缺失最大值更新会导致怎样的错误,并给出修正后的正确更新步骤。

13.1 缺失最大值更新的错误

考虑两个块的计算:块1的最大值为10,块2的最大值为20。如果我们在处理块2时没有正确更新全局最大值,而是直接累加指数(比如始终使用块1的最大值10),则块2的指数可能会因为减去过小的最大值而产生巨大数值,导致溢出;反之如果使用一个过大的固定最大值,则所有指数都非常接近0,导致下溢。数值稳定 softmax 的关键是:每个元素的指数必须减去当前已知的全局最大值,而全局最大值随着处理更多块而单调增加。

错误示例:假定处理块1后 m_global=10,l_global=sum(exp(s1-10))。处理块2时,块2的实际最大值是20。如果我们没有更新 m_global,直接计算 exp(s2-10),那么块2的指数将是 $ \exp(20-10)=\exp(10)\approx22026 $,这本身不会溢出,但问题是当最终归一化时,分母会因为块2的巨大贡献而使得块1的权重几乎为零。这不会造成数值错误(溢出),但另一种情况:如果之前我们使用的最大值过大,比如误用 m=30,那么 $ \exp(20-30)=\exp(-10) $ 非常小,可能导致下溢,所有概率变为0。更重要的是,如果没有更新最大值,后续块可能会改变整个分布,但因为我们没有缩放之前累加的 l 和 o,最终结果在数学上是错误的。

13.2 正确更新步骤(online softmax 归约)

对于每个新块 j,已知其局部最大值 $ m_j $,局部指数和 $ l_j $,局部输出 $ o_j $,全局维护 $ m_g, l_g, o_g $。正确更新为:

$$ 1.m_{n e w}=\max(m_{g},m_{j}) $$

  1. 缩放旧的全局和: $ l_g \leftarrow l_g \cdot e^{m_g - m_{new}} + l_j \cdot e^{m_j - m_{new}} $

  2. 缩放旧的全局输出: $ o_g \leftarrow o_g \cdot e^{m_g - m_{new}} + o_j \cdot e^{m_j - m_{new}} $

  3. $ m_{g} = m_{new} $

13.3 Python 演示错误与修正

代码块

1 import numpy as np

2

3 def wrong_merge(blocks): # 始终使用第一个块的最大值 4 m = blocks[0]['m'] 5 l = blocks[0]['l'] 6 o = blocks[0]['o']


for b in blocks[1]: # 错误:没有更新 m,直接用旧的 m 处理新块 # 这会导致新块的 exp 被错误缩放 l += b['l'] * np.exp(b['m'] - m) # 实际上应该根据新 m 调整 o += b['o'] * np.exp(b['m'] - m) return o / l

def correct_merge(blocks): m = -np.inf l = 0.0 o = 0.0 for b in blocks: m_new = max(m, b['m'] ) l = l * np.exp(m - m_new) + b['l'] * np.exp(b['m'] - m_new) o = o * np.exp(m - m_new) + b['o'] * np.exp(b['m'] - m_new) m = m_new return o / l

用实际数据测试即可观察到错误合并产生的输出与正确合并不一致,且当最大值差异大时错误更明显。

  1. 为 FlashAttention 增加 dropout 支持,前向时应用 dropout mask,并说明反向中如何高效利用保存的 mask 或重计算 mask,编写核心片段。

14.1 前向修改

在计算 softmax 之后,对注意力权重 P 应用 dropout。通常 dropout 在训练时使用,推理时关闭。在前向分块计算中,对每个 (Q块, K块) 对生成 dropout mask(形状 [Br, Bc]),应用 P = P * mask / (1-dropout_prob) (或直接使用 PyTorch 的 dropout 层)。同时需要保存该 mask 或随机种子以便反向使用。

修改 FlashAttention 的内循环(对单个 K/V 块):

代码块

在计算 $ P_{ij} = \exp(S_{ij} - mi_{new}) $ 之后

if self.training and dropout_p > 0: mask = torch.empty_like(P_{ij}).bernoulli_(1 - dropout_p) P_{ij} = $ P_{ij} * mask / (1 - dropout_p) $ else: mask = None

然后 $ li = li + P_{ij}.sum(\ldots) $

以及 $ O_{i} = O_{i} + P_{ij} @ V_{j} $


在实际 kernel 中,需要将 mask 保存到全局内存或通过重计算来避免存储开销。