Block Pointer 与多维支持
会员专享从单序列扩展到 Batch/Head 并行,并使用 Block Pointer 简化指针管理。
在权益中心获取代码在前两章中,我们实现了一个功能完整且经过 Autotune 优化的 Flash Attention Kernel。回顾一下,我们的并行策略是在 序列维度 上进行分块:每个 Program 处理一个 Q Block,然后遍历所有 K/V Block。
但这只利用了一个维度的并行性。实际上,我们可以进一步扩展并行度——真实的 Transformer 模型中,输入张量的形状是 (Batch, Head, SeqLen, Dim),Batch 和 Head 维度天然独立,完全可以并行处理。
本章我们将解决两个问题:
- 多维并行:将并行度从 SeqLen 扩展到 Batch × Head × SeqLen,充分利用 GPU 的并行能力
- 指针管理:用 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_b = 4 × 8 × 64 = 2048— 跳到下一个 batchstride_h = 8 × 64 = 512— 跳到下一个 headstride_m = 64— 跳到下一行(序列维度)stride_d = 1— 跳到下一列(最内层维度)
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这也是为什么 PyTorch 中可以直接用 Q.stride() 获取——它根据 shape 自动计算。
因此,要访问第 (pid_b, pid_h) 个子矩阵,我们需要计算它在一维内存中的偏移:pid_b * stride_b + pid_h * stride_h。这正是下面代码要做的事情。
3D Grid 并行化
在 04_batch_head.py 中,我们使用 3D Grid 来并行化这三个维度:
def call_flash_attention(Q, K, V):
B, H, N, D = Q.shape
O = torch.empty_like(Q)
# 3D Grid: (Batch, Head, SeqBlocks)
grid = lambda META: (B, H, triton.cdiv(N, META["BLOCK_M"]))
flash_attention[grid](
Q, K, V, O,
N, D,
Q.stride(0), Q.stride(1), Q.stride(2), # stride_b, stride_h, stride_m
K.stride(0), K.stride(1), K.stride(2),
V.stride(0), V.stride(1), V.stride(2),
O.stride(0), O.stride(1), O.stride(2),
)
return OGrid 维度映射:
pid_b = tl.program_id(0)→ 第几个 Batchpid_h = tl.program_id(1)→ 第几个 Headpid_m = tl.program_id(2)→ 第几个 Q Block
手动指针偏移
在 Kernel 内部,我们需要手动计算每个 Batch/Head 的基地址偏移:
登录以继续阅读
这是一篇付费内容,请登录您的账户以访问完整内容。
CookLLM文档