编写第一个 Flash Attention Kernel,并利用 Auto-Tune 进行性能优化。
配套代码
我们在上一章推导出了 Flash Attention 的数学公式(Tiling + Online Softmax)。现在,是时候把数学变成代码了。
现在让我们看 Kernel 内部的实现。还记得上一章推导的 “Outer Q, Inner K” 循环逻辑吗?在 Triton 中,它对应如下结构:
systems/flash_attention/01_basic.py
# ... imports ...
@triton.jit
def flash_attention(
q_ptr, k_ptr, v_ptr, o_ptr, # 指针参数
seq_len, d_model: tl.constexpr,
stride_qm, stride_km, stride_vm, stride_om, # Strides
BLOCK_M: tl.constexpr, BLOCK_N: tl.constexpr,
):
# 获取当前 Program 的 ID
pid_m = tl.program_id(0)
start_m = pid_m * BLOCK_M
# [1] 指针计算:生成 Q 的二维访问网格
# offs_m: 逻辑行号, offs_d: 逻辑列号
offs_m = start_m + tl.arange(0, BLOCK_M)
offs_d = tl.arange(0, d_model)
# [2] 初始化 Accumulator
# m_i: running max, l_i: running sum, acc: 结果累加器
m_i = tl.full((BLOCK_M,), float("-inf"), dtype=tl.float32)
l_i = tl.zeros((BLOCK_M,), dtype=tl.float32)
acc = tl.zeros((BLOCK_M, d_model), dtype=tl.float32)
# 加载 Loop-Invariant 的 Q (Q 在内循环中是不变的)
q = tl.load(
q_ptr + offs_m[:, None] * stride_qm + offs_d[None, :],
mask=(offs_m[:, None] < seq_len) & (offs_d[None, :] < d_model),
other=0.0,
)
# [3] 内循环:遍历所有的 K, V 块
for start_n in range(0, seq_len, BLOCK_N):
offs_n = start_n + tl.arange(0, BLOCK_N)
# 动态加载 K, V
k = tl.load(
k_ptr + offs_n[:, None] * stride_km + offs_d[None, :],
mask=(offs_n[:, None] < seq_len) & (offs_d[None, :] < d_model),
other=0.0,
)
v = tl.load(
v_ptr + offs_n[:, None] * stride_vm + offs_d[None, :],
mask=(offs_n[:, None] < seq_len) & (offs_d[None, :] < d_model),
other=0.0,
)
# [4] 计算 Attention Score (QK^T)
qk = tl.dot(q, tl.trans(k)).to(tl.float32)
qk *= 1.0 / tl.sqrt(tl.cast(d_model, tl.float32))
# [5] Online Softmax 更新逻辑 (与上一章推导完全一致!)
m_ij = tl.max(qk, axis=1)
m_next = tl.maximum(m_i, m_ij)
alpha = tl.exp(m_i - m_next)
beta = tl.exp(qk - m_next[:, None])
# 更新 acc (修正旧值 + 加上新贡献)
acc = acc * alpha[:, None] + tl.dot(beta.to(tl.float16), v)
# 更新统计量
l_i = l_i * alpha + tl.sum(beta, axis=1)
m_i = m_next
Triton 的编程范式:
注意我们没有显式写 “Outer Q Loop”。因为在 Triton 的 SPMD 模型中(概念类似 Triton 的 SPMD 编程模型):
def call_flash_attention(Q, K, V):
seq_len, d_model = Q.shape
O = torch.empty_like(Q)
BLOCK_M = 32
BLOCK_N = 32
grid = (triton.cdiv(seq_len, BLOCK_M),)
flash_attention[grid](
Q,
K,
V,
O,
seq_len,
d_model,
Q.stride(0),
K.stride(0),
V.stride(0),
O.stride(0),
BLOCK_M=BLOCK_M,
BLOCK_N=BLOCK_N,
)
return O
这就是为什么代码里看起来只有一个 Inner Loop,却能并行处理整个序列。
你可能注意到了 d_model: tl.constexpr 和 BLOCK_M: tl.constexpr 这样的标注。为什么它们是必需的?
Triton 的编译期常量要求:
offs_m = start_m + tl.arange(0, BLOCK_M) # BLOCK_M 必须是 constexpr
offs_d = tl.arange(0, d_model) # d_model 必须是 constexpr
tl.arange() 函数要求其参数在编译时已知,因为: