编写第一个 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,却能并行处理整个序列。

理解 tl.constexpr 的必要性

你可能注意到了 d_model: tl.constexprBLOCK_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() 函数要求其参数在编译时已知,因为: