先读调用契约:把每个 Block 的前向交给 checkpoint
🎯 先猜一猜
循环已经来到第 3 个 Transformer Block,当前变量 x 是前两个 Block 更新后的状态。
哪种写法既执行当前 Block,又把结果继续传给下一层?
先补的知识
- blocks 是 nn.ModuleList;循环中的 block 是一个可调用的 nn.Module,调用 block(x) 会返回下一层输入。
- x 是需要梯度的浮点张量。Checkpoint 不改变它的逻辑 shape;本题的测试输入是 [2, 2048, 2048],device 是 CUDA。
- checkpoint 的第一个位置参数是要执行的函数,后面的位置参数是传给该函数的输入;use_reentrant=False 是本 Notebook 要求的关键字参数。
- 循环必须把返回值重新赋给 x,否则后面的 Block 仍会收到旧状态,网络就没有按层向前传播。
图解原理
普通循环写的是 x = block(x)。Checkpoint 版本没有换模型,也没有跳过计算,只是在 block 外面加了一个“反向时可以重算”的包装器。因此最稳的解题方式是先写出普通调用,再把 block(x) 改写成 checkpoint(block, x, ...)。
普通版本
Autograd 为反向传播保留 Block 内部需要的中间激活。
Checkpoint 版本
数学前向不变,变化的是 Autograd 保存与重算中间量的策略。
语法热身:用一个不同的小函数练“函数和参数分开传入”
import torch
from torch.utils.checkpoint import checkpoint
def add_bias(values, bias):
return values + bias
values = torch.randn(4, requires_grad=True)
bias = torch.ones(4, requires_grad=True)
result = checkpoint(add_bias, values, bias, use_reentrant=False)
result.sum().backward()从语法例子迁移到 Notebook
add_bias:对应循环中的 block;二者都是可以稍后再次调用的 callablevalues:对应当前 x;它是传给 callable 的输入张量result:对应 checkpoint 返回的新 x,必须接住并继续向后传use_reentrant=False:原样映射到 TODO,不能误当成 block 的参数
巩固一下
为什么 TODO 中不能只写 checkpoint(block, x, use_reentrant=False) 而不赋值?
学完这一段,试着做
用一个小动作确认自己理解了;最后再进入官方题目。