从单序列扩展到 Batch/Head 并行,并使用 Block Pointer 简化指针管理。

配套代码


在前两章中,我们实现了一个功能完整且经过 Autotune 优化的 Flash Attention Kernel。回顾一下,我们的并行策略是在 序列维度 上进行分块:每个 Program 处理一个 Q Block,然后遍历所有 K/V Block。

但这只利用了一个维度的并行性。实际上,我们可以进一步扩展并行度——真实的 Transformer 模型中,输入张量的形状是 (Batch, Head, SeqLen, Dim)Batch 和 Head 维度天然独立,完全可以并行处理。

本章我们将解决两个问题:

  1. 多维并行:将并行度从 SeqLen 扩展到 Batch × Head × SeqLen,充分利用 GPU 的并行能力
  2. 指针管理:用 Block Pointer 简化 4D 张量带来的复杂地址计算

从单序列到 Batch/Head 并行

4D 张量的内存布局

当输入从 (N, D) 变为 (B, H, N, D) 时,指针计算变得复杂。虽然我们逻辑上看到的是一个 4D 张量,但 GPU 显存中它仍然是一维连续存储的(回顾 张量布局 中的 Stride 概念)。

以一个形状为 (2, 4, 8, 64) 的张量为例(2 个 batch,4 个 head,序列长度 8,维度 64):

逻辑视图: Q[batch, head, seq, dim]  →  Q[2, 4, 8, 64]
物理存储: 一维数组,共 2 × 4 × 8 × 64 = 4096 个元素

Stride 告诉我们在每个维度上移动一个单位需要跳过多少个元素:

Stride 的计算规律:对于行主序存储,每个维度的 stride = 后面所有维度大小的乘积。最后一个维度的 stride 总是 1,因为相邻元素在物理内存中也相邻。

Shape:   (B,     H,     N,    D )
         (2,     4,     8,    64)
          ↓      ↓      ↓     ↓
Stride:  H×N×D   N×D    D     1
         2048    512    64    1