17. Autograd Basics | Attention 反向传播与自定义 Autograd
难度: Medium | 环境: CPU-first | 标签: 显存优化, Autograd, 反向传播 | 目标人群: 显存优化学习者
🚀 云端运行环境
本章节的实战代码可以点击以下链接在免费 GPU 算力平台上直接运行:
本节导读
在掌握 Part00 07 节的计算图、backward() 和 autograd.Function 基础后,再来看 Attention 的反向路径:损失的梯度到底怎样回到 loss.backward(),这条路径会被 PyTorch 自动处理;但一旦要写自定义算子、排查梯度异常,或者理解后面 FlashAttention 的反向传播,就需要把这条链路拆开看。
本节不重复 Autograd 的通用定义,而是沿着一个简化版 Attention 反向走一遍:输出梯度先回到 torch.autograd.Function,理解为什么前向要保存中间张量,并用 PyTorch 自动求导结果校验手写梯度是否正确。
本节的代码任务是完成简化 Attention backward,并用 reference 和 gradcheck 验证:输出梯度沿矩阵乘法和 Softmax 回传到 Q / K / V。范围限于无 mask、无 dropout 的 CPU-first 实现;ctx.save_for_backward 只用于说明本题需要的状态,不代表真实 GPU 显存、带宽、GQA/MQA、KV Cache 或 FlashAttention 性能结论。
关键词: Autograd, backward, gradcheck
前置阅读
导语: 先理解 PyTorch 如何记录计算图、如何执行训练循环,再进入手写 Attention backward 会更顺。
- P0: 07. PyTorch Autograd and Backward | PyTorch 自动求导与反向传播
- P0: 13. Simple Neural Network Training | 简单神经网络训练循环
相关阅读
导语: 手写 Attention backward 之后,可以继续看梯度如何穿过激活函数、损失函数,以及性能分析里如何定位反向算子的开销。
Step 1: 前向传播回顾与变量定义
为了和后面的代码保持一致,我们先回顾 04 节的缩放点积 Attention 前向公式:
- 打分矩阵:
- 概率矩阵:
- 最终输出:
张量形状说明:
(batch 、序列长度 、特征维数 )
Step 2: 链式法则
假设下游的损失函数已经帮我们算好了输出张量
1. 求
2. 求
3. 计算 Softmax 的梯度 我们需要从
(注:
4. 求
Step 3: 手撕 PyTorch Autograd Function
现在,把你刚才看到的微积分公式,转化为能够实际运行的代码。我们将继承 torch.autograd.Function。
要求:完成 backward 函数中 TODO 的数学推导代码。你可以使用 ctx.saved_tensors 来获取前向传播时保存的
本题只实现无 mask、无 dropout 的缩放点积 Attention backward;causal mask、GQA/MQA 和 KV Cache 不在本题范围内。 这一节的实现顺序就是先求 dV / dP,再穿过 Softmax 得到 dS,最后回到 dQ / dK。
提示
- 反向顺序可以记成
dV -> dP -> dS -> dQ / dK,先把最容易的分支算出来,再往 Softmax 和输入侧回推。 ctx.save_for_backward里保存的是本题手写 backward 所需的状态;生产实现可能通过重算、分块或其他数据流减少保存量。- Softmax 的反向不要去构造完整雅可比,按行做修正项即可。
import torch
import torch.nn.functional as F
import mathclass CustomAttention(torch.autograd.Function):
"""无 mask、无 dropout 的教学版缩放点积 Attention。
输入和输出均使用 [B, N, d];本题只实现单头、等长 Q/K/V 的路径。
"""
@staticmethod
def forward(ctx, q, k, v):
"""计算缩放点积 Attention。
Args:
q, k, v: [B, N, d] 张量,device 和 dtype 必须一致。
Returns:
[B, N, d] 的 attention 输出。
"""
if q.ndim != 3 or k.ndim != 3 or v.ndim != 3:
raise ValueError('q、k、v 必须是 [batch, seq_len, head_dim] 三维张量')
if q.shape[:2] != k.shape[:2] or k.shape[:2] != v.shape[:2]:
raise ValueError('q、k、v 的 batch 和序列长度必须一致')
if q.size(-1) != k.size(-1) or v.size(-1) != q.size(-1):
raise ValueError('q、k、v 的 head_dim 必须一致')
if q.device != k.device or k.device != v.device:
raise ValueError('q、k、v 必须位于同一 device')
if q.dtype != k.dtype or k.dtype != v.dtype:
raise TypeError('q、k、v 必须使用相同 dtype')
# 1. 缩放点积
d_k = q.size(-1)
scale = 1.0 / math.sqrt(d_k)
scores = torch.matmul(q, k.transpose(-2, -1)) * scale
# 2. Softmax 获取概率 P
p = F.softmax(scores, dim=-1)
# 3. 乘上 V 得到输出
out = torch.matmul(p, v)
# 保存反向传播需要用到的张量
ctx.save_for_backward(q, k, v, p)
ctx.scale = scale
return out
@staticmethod
def backward(ctx, dout):
"""根据上游输出梯度返回 q、k、v 的梯度。
Args:
dout: 与 forward 输出同形状的上游梯度。
Returns:
dq、dk、dv,形状分别与 q、k、v 相同。
"""
# 提取前向保存的张量
q, k, v, p = ctx.saved_tensors
scale = ctx.scale
# ==========================================
# TODO 1: 求 dV
# 提示:由 out = P @ V,使用 P.transpose(-2, -1) @ dout。
# ==========================================
# dv = ???
# ==========================================
# TODO 2: 求 dP
# 提示:由 out = P @ V,使用 dout @ V.transpose(-2, -1)。
# ==========================================
# dp = ???
# ==========================================
# TODO 3: 穿过 Softmax 求 dS
# 提示:对最后一维求 row_sum,避免显式构造 Softmax 雅可比矩阵。
# ==========================================
# dp_mul_p = ???
# row_sum = ???
# ds = ???
# ==========================================
# TODO 4: 求 dQ 和 dK (别忘了乘以 scale 缩放因子)
# 提示:由 S = Q @ K.transpose(-2, -1) * scale 回传到 q / k。
# ==========================================
# dq = ???
# dk = ???
return dq, dk, dv测试
运行下面的测试单元,确认手写 backward 和 PyTorch 自动求导保持一致。
# 运行此单元格以测试你的实现
def test_attention_backward():
try:
torch.manual_seed(42)
B, N, d = 2, 8, 16
# 随机初始化张量,必须要求梯度
q = torch.randn(B, N, d, dtype=torch.float64, requires_grad=True)
k = torch.randn(B, N, d, dtype=torch.float64, requires_grad=True)
v = torch.randn(B, N, d, dtype=torch.float64, requires_grad=True)
print("1. 测试前向传播是否正常...")
custom_out = CustomAttention.apply(q, k, v)
# 原生 PyTorch 实现
scores = torch.matmul(q, k.transpose(-2, -1)) / math.sqrt(d)
ref_out = torch.matmul(F.softmax(scores, dim=-1), v)
assert torch.allclose(custom_out, ref_out), "前向传播结果不一致!"
assert custom_out.shape == (B, N, d), "Attention 输出形状不一致!"
print("\n2. 进行梯度数值检验 (Gradcheck)...")
test_passed = torch.autograd.gradcheck(CustomAttention.apply, (q, k, v), eps=1e-6, atol=1e-4)
# 用同一组输入分别回传,检查 q/k/v 的梯度形状和数值。
q_ref, k_ref, v_ref = [t.detach().clone().requires_grad_() for t in (q, k, v)]
ref_scores = torch.matmul(q_ref, k_ref.transpose(-2, -1)) / math.sqrt(d)
ref_output = torch.matmul(F.softmax(ref_scores, dim=-1), v_ref)
ref_output.sum().backward()
q_custom, k_custom, v_custom = [t.detach().clone().requires_grad_() for t in (q, k, v)]
CustomAttention.apply(q_custom, k_custom, v_custom).sum().backward()
assert q_custom.grad.shape == q_custom.shape
assert k_custom.grad.shape == k_custom.shape
assert v_custom.grad.shape == v_custom.shape
assert torch.allclose(q_custom.grad, q_ref.grad, atol=1e-5)
assert torch.allclose(k_custom.grad, k_ref.grad, atol=1e-5)
assert torch.allclose(v_custom.grad, v_ref.grad, atol=1e-5)
if test_passed:
print("✅ All Tests Passed! Attention 反向传播实现通过测试。")
except NotImplementedError:
print("请先完成 TODO 部分的代码!")
raise
except (AttributeError, NameError, TypeError, ValueError, AssertionError, RuntimeError) as e:
if isinstance(e, AttributeError):
print("代码未完成,无法找到必要的属性")
elif isinstance(e, NameError):
print("代码可能未完成,导致了变量未定义")
elif isinstance(e, TypeError):
print("代码可能未完成,导致了类型错误")
elif isinstance(e, ValueError):
print("代码可能未完成,导致了张量维度错误")
elif isinstance(e, AssertionError):
print(f"代码可能未完成,导致了断言失败: {e}")
else:
print("代码可能未完成,导致了 gradcheck 或反向传播异常")
raise NotImplementedError("请先完成 TODO 部分的代码!") from e
except Exception as e:
print(f"❌ 发生异常: {e}")
raise
test_attention_backward()Step 4: 从显式 Attention 到分块实现(预告)
看看你刚才写的 ctx.save_for_backward(q, k, v, p)。这行代码在反向传播被调用前,会一直把
如果序列长度为
思考题:怎样减少这个中间矩阵的驻留和搬运? 提示:FlashAttention 通过分块计算、在线 Softmax 和片上数据复用,避免显式保存完整的
。实际显存和速度收益仍需结合 GPU、序列长度、dtype 和 workload 测量。
下一节将用一个简化实现说明 FlashAttention 如何改变这条数据流。
🛑 STOP HERE 🛑
请先尝试自己完成代码并跑通测试。
如果你正在 Colab 中运行,并且遇到困难没有思路,可以向下滚动查看参考答案。
参考代码与解析
代码
class CustomAttention(torch.autograd.Function):
"""无 mask、无 dropout 的教学版缩放点积 Attention。
输入和输出均使用 [B, N, d];本题只实现单头、等长 Q/K/V 的路径。
"""
@staticmethod
def forward(ctx, q, k, v):
"""计算缩放点积 Attention。
Args:
q, k, v: [B, N, d] 张量,device 和 dtype 必须一致。
Returns:
[B, N, d] 的 attention 输出。
"""
if q.ndim != 3 or k.ndim != 3 or v.ndim != 3:
raise ValueError('q、k、v 必须是 [batch, seq_len, head_dim] 三维张量')
if q.shape[:2] != k.shape[:2] or k.shape[:2] != v.shape[:2]:
raise ValueError('q、k、v 的 batch 和序列长度必须一致')
if q.size(-1) != k.size(-1) or v.size(-1) != q.size(-1):
raise ValueError('q、k、v 的 head_dim 必须一致')
if q.device != k.device or k.device != v.device:
raise ValueError('q、k、v 必须位于同一 device')
if q.dtype != k.dtype or k.dtype != v.dtype:
raise TypeError('q、k、v 必须使用相同 dtype')
d_k = q.size(-1)
scale = 1.0 / math.sqrt(d_k)
scores = torch.matmul(q, k.transpose(-2, -1)) * scale
p = F.softmax(scores, dim=-1)
out = torch.matmul(p, v)
ctx.save_for_backward(q, k, v, p)
ctx.scale = scale
return out
@staticmethod
def backward(ctx, dout):
"""根据上游输出梯度返回 q、k、v 的梯度。
Args:
dout: 与 forward 输出同形状的上游梯度。
Returns:
dq、dk、dv,形状分别与 q、k、v 相同。
"""
q, k, v, p = ctx.saved_tensors
scale = ctx.scale
# TODO 1: 求 dV
# 提示:由 out = P @ V,使用 P.transpose(-2, -1) @ dout。
dv = torch.matmul(p.transpose(-2, -1), dout)
# TODO 2: 求 dP
# 提示:由 out = P @ V,使用 dout @ V.transpose(-2, -1)。
dp = torch.matmul(dout, v.transpose(-2, -1))
# TODO 3: 穿过 Softmax 求 dS
# 提示:对最后一维求 row_sum,避免显式构造 Softmax 雅可比矩阵。
dp_mul_p = dp * p
row_sum = dp_mul_p.sum(dim=-1, keepdim=True)
ds = p * (dp - row_sum)
# TODO 4: 求 dQ 和 dK
# 提示:由 S = Q @ K.transpose(-2, -1) * scale 回传到 q / k。
dq = torch.matmul(ds, k) * scale
dk = torch.matmul(ds.transpose(-2, -1), q) * scale
return dq, dk, dv解析
1. TODO 1: 求 dV
- 实现方式:
dv = torch.matmul(p.transpose(-2, -1), dout) - 数学原理:输出
out = P V,所以对V的梯度就是把上游梯度乘回去。 - 工程意义:这是 Attention 反向里最直接的一步,先把 Value 方向的梯度拿到。
2. TODO 2: 求 dP
- 实现方式:
dp = torch.matmul(dout, v.transpose(-2, -1)) - 数学原理:同样由
out = P V得到,对P的梯度可以直接通过矩阵乘法反推。 - 工程意义:这一步把输出梯度重新映射回注意力概率矩阵。
3. TODO 3: 穿过 Softmax 求 dS
- 实现方式:先算
dp_mul_p = dp * p,再对行求和得到row_sum,最后得到ds = p * (dp - row_sum)。 - 数学原理:Softmax 的反向可以化成一个稳定的逐行修正项,不需要显式构造完整雅可比矩阵。
- 工程意义:这一步决定了 Softmax 输出梯度如何回到打分矩阵,是实现中需要重点检查的环节。
4. TODO 4: 求 dQ 和 dK
- 实现方式:
dq = torch.matmul(ds, k) * scale,dk = torch.matmul(ds.transpose(-2, -1), q) * scale - 数学原理:因为
scores = QK^T / sqrt(d),所以最后回到Q和K时还要乘上缩放因子。 - 工程意义:这一步把 Attention 的反向梯度真正落回输入表示。
进阶思考
- 如果不保存
P,反向传播还能怎么做? - 为什么这会自然引出 FlashAttention 的重计算思想?
- 你能把这条链路和第 20 节的在线 Softmax 对上吗?
