从单序列扩展到 Batch/Head 并行,并使用 Block Pointer 简化指针管理。
配套代码
在前两章中,我们实现了一个功能完整且经过 Autotune 优化的 Flash Attention Kernel。回顾一下,我们的并行策略是在 序列维度 上进行分块:每个 Program 处理一个 Q Block,然后遍历所有 K/V Block。
但这只利用了一个维度的并行性。实际上,我们可以进一步扩展并行度——真实的 Transformer 模型中,输入张量的形状是 (Batch, Head, SeqLen, Dim),Batch 和 Head 维度天然独立,完全可以并行处理。
本章我们将解决两个问题:
当输入从 (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