LogoCookLLM文档
LogoCookLLM文档
首页CookLLM

原理精讲

词元化
Tokenization 基础BPE 算法详解GPT 系列 TokenizerBPE 训练工程化
模型架构
Transformer LM
从 token ids 到 logitsEmbedding 与 LM Head
Attention 机制
Self-Attention 到 GQAAttention Sink
位置编码
位置编码基础RoPE 数学推导RoPE 代码实现长度外推
GPU 编程基础
GPU 架构基础张量布局Triton 入门:向量加法
FlashAttention
Flash Attention 原理详解从朴素实现到 Auto-TuningBlock Pointer 与多维支持Causal Masking 优化Grouped Query Attention反向传播实现
分布式训练
数据并行ZeRO 优化器全分片数据并行张量并行流水线并行多维混合并行
推理优化
KV CacheContinuous BatchingPagedAttention

动手训练

概述
预训练
预训练数据Tokenizer 训练模型架构数据流水线训练循环监控与验证
X (Twitter)
系统工程FlashAttention

Block Pointer 与多维支持

会员专享

从单序列扩展到 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_b = 4 × 8 × 64 = 2048 — 跳到下一个 batch
  • stride_h = 8 × 64 = 512 — 跳到下一个 head
  • stride_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 来并行化这三个维度:

systems/flash_attention/04_batch_head.py
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 O

Grid 维度映射:

  • pid_b = tl.program_id(0) → 第几个 Batch
  • pid_h = tl.program_id(1) → 第几个 Head
  • pid_m = tl.program_id(2) → 第几个 Q Block

手动指针偏移

在 Kernel 内部,我们需要手动计算每个 Batch/Head 的基地址偏移:

登录以继续阅读

这是一篇付费内容,请登录您的账户以访问完整内容。

从朴素实现到 Auto-Tuning

编写第一个 Flash Attention Kernel,并利用 Auto-Tune 进行性能优化。

Causal Masking 优化

为自回归模型实现因果注意力机制,通过跳过上三角计算实现 ~2x 加速。

目录

从单序列到 Batch/Head 并行
4D 张量的内存布局
3D Grid 并行化
手动指针偏移
Block Pointer:优雅的解决方案
核心 API
循环中的指针移动
完整对比
总结