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

反向传播实现

会员专享

实现 Flash Attention 的梯度计算,通过 Recomputation 实现内存高效的训练。

在权益中心获取代码

在前面的章节中,我们实现了完整的 Flash Attention 前向传播,支持任意序列长度、Causal Masking 和 GQA。但这还不够——要在训练中使用 Flash Attention,我们需要实现反向传播(Backward Pass)来计算梯度。

本章将探讨如何在保持 IO 效率的同时计算 ∂L∂Q\frac{\partial L}{\partial Q}∂Q∂L​、∂L∂K\frac{\partial L}{\partial K}∂K∂L​、∂L∂V\frac{\partial L}{\partial V}∂V∂L​。

为什么需要自定义反向传播?

PyTorch 自动微分的局限性

PyTorch 的 autograd 可以自动为大部分操作生成反向传播代码,但对于 Flash Attention 这样的融合 Kernel,自动微分会遇到问题:

# 标准 Attention 的前向传播
def standard_attention(Q, K, V):
    S = Q @ K.T / sqrt(d)     # (1) 计算分数
    P = softmax(S, dim=-1)    # (2) Softmax
    O = P @ V                 # (3) 加权求和
    return O

自动微分的行为:

  • PyTorch 会保存中间结果 SSS 和 PPP (Attention Matrix)
  • 在反向传播时使用这些中间结果计算梯度
  • 内存开销: O(N2)O(N^2)O(N2) 用于存储 Attention Matrix

但 Flash Attention 的核心优势就是不物化(materialize) Attention Matrix!

如果使用自动微分,我们会:

  1. 前向传播:不保存 Attention Matrix,内存 O(N)O(N)O(N) ✅
  2. 反向传播:需要 Attention Matrix,被迫重新计算,失去了优势 ❌

Recomputation 策略

Flash Attention 采用了一个巧妙的权衡:Recomputation(重计算)

核心思想:

  • 前向传播时:只保存少量的中间统计量 (L\mathbf{L}L 和 M\mathbf{M}M)
  • 反向传播时:使用这些统计量重新计算 Attention 分数,而非从 HBM 加载
传统方法 (Materialization):
  Forward:  计算 S, P → 保存到 HBM (O(N²) 内存)
  Backward: 从 HBM 读取 S, P → 计算梯度

Flash Attention (Recomputation):
  Forward:  计算 S, P → 只保存 L, M (O(N) 内存)
  Backward: 重新计算 S, P → 计算梯度

权衡分析:

  • ✅ 内存节省:O(N2)→O(N)O(N^2) \to O(N)O(N2)→O(N)
  • ⚠️ 额外计算:需要重新计算 Attention(约 1.5x FLOPS)
  • ✅ IO 效率:仍然比标准实现快,因为避免了 HBM 往返

为什么 Recomputation 仍然更快?

虽然增加了 FLOPS,但现代 GPU 是 IO-bound 而非 compute-bound。重新计算 Attention 的 FLOPS 开销小于从 HBM 加载 O(N2)O(N^2)O(N2) 矩阵的 IO 开销。

这就是为什么 Flash Attention 即使在反向传播中也能保持性能优势。

登录以继续阅读

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

Grouped Query Attention

实现 GQA/MQA 支持,让多个 Query Head 共享 KV,优化 KV Cache 内存占用。

分布式训练

从数据并行到多维混合并行,理解大模型训练的核心并行策略

目录

为什么需要自定义反向传播?
PyTorch 自动微分的局限性
Recomputation 策略
Attention 反向传播的数学原理
前向传播回顾
梯度推导
1. ∂L∂V\frac{\partial \mathcal{L}}{\partial \mathbf{V}}∂V∂L​ (最简单)
2. ∂L∂P\frac{\partial \mathcal{L}}{\partial \mathbf{P}}∂P∂L​ (中间梯度)
3. ∂L∂S\frac{\partial \mathcal{L}}{\partial \mathbf{S}}∂S∂L​ (Softmax 反向传播)
4. ∂L∂Q\frac{\partial \mathcal{L}}{\partial \mathbf{Q}}∂Q∂L​ 和 ∂L∂K\frac{\partial \mathcal{L}}{\partial \mathbf{K}}∂K∂L​
完整的梯度计算流程
代码实现解析
Forward Kernel 的修改
Backward Kernel 实现
关键实现细节
1. 循环顺序的反转
2. Atomic Add 处理 dQ
3. 重新计算 P 而非保存
torch.autograd.Function 封装
性能验证
数值正确性测试
内存占用对比
设计权衡与优化方向
Recomputation 的开销
进一步优化方向
总结