返回关卡地图
预判 0 / 3 闯关题:0 / 3 XP 0 / 300

Level 17 | 零基础导学关卡

Attention 反向传播与自定义 Autograd

Attention Backward and Custom Autograd

先通过图解和小例子理解原理,再做预测与练习;不懂的地方可以反复尝试,最后把理解写进官方 Notebook。

本课官方 Notebook ↗
17_Autograd_Basics.ipynb 显存优化Autograd反向传播
Mission 1 概念练习

从 out=P@V 分出 dV 与 dP 两条梯度支路

🎯 先猜一猜

p.shape=[2,8,8]、v.shape=[2,8,16]、dout.shape=[2,8,16]。

dv 与 dp 的 shape 分别是什么?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • 前向中 q/k/v.shape=[B,N,d],scores/p.shape=[B,N,N],out/dout.shape=[B,N,d]。
  • 矩阵乘法反向会使用转置:out=P@V,所以 dV=P^T@dout,dP=dout@V^T。
  • transpose(-2,-1) 只交换最后两个矩阵维,保留 batch 维。
  • ctx.saved_tensors 按 forward 保存顺序取回 q、k、v、p;它们的 dtype/device 与前向一致。

图解原理

out 的每个值同时依赖注意力概率 P 和内容 V。上游梯度 dout 到达矩阵乘法后分成两路:一路告诉 V 应怎样变化,另一路告诉 P 应怎样变化;之后 P 的梯度还要继续穿过 Softmax。

out = P @ Vdout [B,N,d] P 支路:dP=dout@VᵀV 支路:dV=Pᵀ@dout
计算左输入右输入结果
dV=Pᵀ@dout[B,N,N][B,N,d][B,N,d]
dP=dout@Vᵀ[B,N,d][B,d,N][B,N,N]
不要使用 .T;对三维 Tensor,.T 的行为不等同于只转置矩阵最后两维。这里明确使用 transpose(-2,-1)。

语法热身:手算批量矩阵乘法的两条反向支路

weights = torch.randn(3, 4, 5)
values = torch.randn(3, 5, 2)
grad_output = torch.randn(3, 4, 2)

grad_values = weights.transpose(-2, -1) @ grad_output
grad_weights = grad_output @ values.transpose(-2, -1)
print(grad_values.shape, grad_weights.shape)

把独立例子迁移回 Notebook

  • weights / values / grad_output 对应 p / v / dout
  • grad_values 对应 dv,shape 跟 v 一致。
  • grad_weights 对应 dp,shape 跟 p 一致。
  • 例子故意使用非方形矩阵,帮助你从维度而不是死记位置判断转置。

巩固一下

dP 为什么使用 dout @ v.transpose(-2,-1)?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

Mission 2 概念练习

穿过 Softmax:用逐行修正替代完整雅可比

🎯 先猜一猜

dp_mul_p.shape=[2,8,8]。

沿 dim=-1 求和并 keepdim=True 后,row_sum.shape 是什么?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • Softmax 沿 dim=-1 逐行计算,因此反向修正项也必须沿最后一维求和。
  • dp 与 p 都是 [B,N,N];逐元素乘法 dp*p 不改变 shape。
  • sum(dim=-1,keepdim=True) 得到 [B,N,1],可广播回 [B,N,N]。
  • 公式是 ds=p*(dp-row_sum),括号与乘法顺序不能随意改变。

图解原理

Softmax 一行中的概率相互耦合:提高一个分数会挤压同一行其他概率。row_sum 汇总这行共同的影响,再从每个 dp 中减去,最后乘 p,得到真正回到 score 的梯度。

1dp_mul_p = dp * p先计算每个位置的加权上游梯度
2row_sum = ...sum(-1, keepdim=True)得到每行一个共享修正量
3ds = p * (dp - row_sum)利用广播把修正量应用到整行

keepdim=True

row_sum 为 [B,N,1],可自然广播到每一列。

漏掉 keepdim

row_sum 变成 [B,N],在一般 batch/sequence shape 下可能广播错轴或直接报错。

语法热身:对一批类别概率手写 Softmax 向量积

prob = torch.softmax(torch.randn(2, 5), dim=-1)
grad_prob = torch.randn(2, 5)

weighted = grad_prob * prob
shared = weighted.sum(dim=-1, keepdim=True)
grad_score = prob * (grad_prob - shared)
print(grad_score.shape)

把独立例子迁移回 Notebook

  • prob / grad_prob 对应 p / dp
  • weighted / shared / grad_score 对应 dp_mul_p / row_sum / ds
  • 独立例子是 [2,5],Notebook 是 [B,N,N];只要始终沿最后一维,公式可直接泛化。
  • keepdim=True 是让 shared 正确广播的语法关键。

巩固一下

Softmax backward 为什么不直接写 ds=dp*p?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

Mission 3 概念练习

回到 Q/K:补回 scale,并让 gradcheck 做最终裁判

🎯 先猜一猜

dq、dk 的矩阵乘法都正确,但漏乘了 scale。

forward allclose 与 gradcheck 最可能怎样?

先猜一个答案,再对照下面的讲解。可以随时修改选择,猜错会得到解释。

先补的知识

  • scores=Q@K^T×scale,scale=1/sqrt(d) 已保存在 ctx.scale。
  • 由矩阵乘法可得 dq=ds@k×scale,dk=ds^T@q×scale。
  • CustomAttention.forward 的输入顺序是 q、k、v,所以 backward 必须返回 dq、dk、dv。
  • gradcheck 使用 float64、eps=1e-6、atol=1e-4,将解析梯度与有限差分数值梯度比较。

图解原理

dS 已经到达打分矩阵,最后沿 QK^T 的两条支路回传。前向把 score 缩小了 scale,链式法则要求两条输入梯度也各乘一次同样的 scale;返回顺序则决定 Autograd 把梯度交给谁。

ds [B,N,N]ds@k×scale → dq [B,N,d] dsᵀ [B,N,N]dsᵀ@q×scale → dk [B,N,d]
前向验收custom_out 与原生 attention ref_out allclose 反向验收torch.autograd.gradcheck(CustomAttention.apply,...) 测试输入B=2、N=8、d=16、dtype=float64、requires_grad=True 返回契约return dq, dk, dv
forward allclose 通过不代表 backward 正确。漏 scale、转置错或返回顺序错,都只能由 gradcheck 揭示。

语法热身:为缩放双线性打分手写两侧梯度

left = torch.randn(2, 3, 4)
right = torch.randn(2, 5, 4)
grad_score = torch.randn(2, 3, 5)
factor = 0.5

grad_left = grad_score @ right * factor
grad_right = grad_score.transpose(-2, -1) @ left * factor

把独立例子迁移回 Notebook

  • left / right 对应 q / k
  • grad_score 对应 dsfactor 对应 scale
  • grad_left / grad_right 对应 dq / dk
  • Notebook 最后还要把前一任务的 dv 按 q、k、v 顺序一起返回。

巩固一下

正确的 backward 返回顺序是什么?

学完这一段,试着做

用一个小动作确认自己理解了;最后再进入官方题目。

把理解变成自己的代码

准备好,去官方题目试一试

在本页用小例子建立直觉,再去官方 Notebook 完成实现。先运行你自己的测试,遇到困难时再查看官方提示与参考答案。

前往本课官方 Notebook ↗

这节课会遇到的代码对象

CustomAttention · forward · backward

先找题目里的输入、输出与 TODO,再把本课的手算过程对应进去;以官方题目中的函数说明和测试为准。

练习来源:Datawhale 官方仓库 · 518cc45。这里的入口直接打开官方版本,不读取或分享你的本地 Notebook。

完成 3 道闯关题后,本关即算完成;作业 checklist 用来辅助你回 notebook 练习。
进度只保存在当前浏览器 localStorage,分享 HTML 不会带走你的记录。
上一关:L16 下一关:L18