为什么你的 PyTorch 代码“能跑”却“不对”
为什么你的 PyTorch 代码“能跑”却“不对”
你可能写过这样的代码:前向传播一切正常,loss 也在下降,但某个关键参数始终不更新;或者 loss.backward() 突然抛出 RuntimeError,提示你“trying to backward through the graph a second time”;又或者训练几个 epoch 后显存悄悄涨满,而你明明已经用了 no_grad。
这些问题的根源,往往不在模型结构或数据加载器里,而在你对计算图的心智模型上。PyTorch 的动态图机制给了你极大的灵活性,但也把调试的责任完全交给了你。本文不讲拓扑排序的数学证明,也不带你从零实现自动微分引擎,我们只聊一件事:作为框架使用者,如何建立正确的计算图直觉,并在报错时快速定位到图结构层面的问题。
张量不只是数组:它携带了整张计算图的线索
在 PyTorch 中,张量从来不是孤立的数据块。当你执行 y = x * 2 + 1 时,y 不仅存储了数值结果,还悄悄记住了一件事:“我是由 MulBackward0 和 AddBackward0 两个操作生成的,我的父节点是 x*2 的结果和常量 1。”
下面这段代码展示了这种“记忆”是如何被编码的:
import torch
x = torch.tensor([2.0], requires_grad=True)
w = torch.tensor([3.0], requires_grad=True)
b = torch.tensor([1.0]) # 注意:b 没有 requires_grad
y = w * x + b
loss = y.sum()
print(f"loss.grad_fn: {loss.grad_fn}") # SumBackward0
print(f"y.grad_fn: {y.grad_fn}") # AddBackward0
print(f"x.is_leaf: {x.is_leaf}, x.grad_fn: {x.grad_fn}") # True, None
print(f"b.requires_grad: {b.requires_grad}") # False
这里有三个关键观察点。首先,grad_fn 是图的边,每个非叶子张量的 grad_fn 指向生成它的算子,该算子又持有对输入张量的引用,这就是反向传播的路径。其次,叶子节点是图的起点,只有 requires_grad=True 且由用户直接创建的张量才是叶子节点,它们的 grad_fn 为 None,梯度最终累积到 .grad 属性上。最后,不参与梯度的张量不保存激活值。b 虽参与前向计算,但因 requires_grad=False,PyTorch 不会为其创建 grad_fn,更重要的是,不会为计算它的梯度而保存上游中间激活值。例如一个冻结的 Linear 层,其输入激活值无需保存用于权重梯度计算;若偏置也被冻结,则输入激活值完全不被保留。这才是冻结参数显著省显存的根本原因,grad_fn 元数据本身的开销其实微不足道。