计算图与梯度累积:为什么训练前必须清空梯度

你写了 optimizer.zero_grad() 却不知道为什么要写。这一篇把 PyTorch 的动态计算图、叶子节点和梯度累积机制讲透,让你彻底明白"清梯度"不是仪式,是物理必然。

学习路线:PyTorch 深度学习基础(43~58) · 第 45 课
第 45 课第 45 课计算图与梯度累积视频封面,展示 loss、backward、grad 和 zero_grad 的训练链路。
Learning Path

PyTorch 深度学习基础(43~58)

在前 20 篇 PyTorch 入门上继续深化,补齐张量操作、自动微分、优化器家族和训练排错能力。

学完本阶段你能做到:独立搭建一个 PyTorch 训练脚本,正确处理 Shape、device、梯度清空、优化器选择和学习率策略,并能根据 loss/accuracy 曲线诊断训练问题。

推荐读法:当前更新主线,建议跟更;每篇都要配合代码自检。

查看完整阶段 · 16 篇
Lesson Guide

这一课怎么学

先看目标,再带着问题读正文。读完后用练习确认自己真的理解了。

本节你会学到

  • 说清楚「计算图与梯度累积:为什么训练前必须清空梯度」解决的核心问题
  • 知道它在「PyTorch 深度学习基础(43~58)」中的位置
  • 把 DataLoader、Linear 层、Tensor 这些关键词联系起来

前置知识

  • 建议先读完上一篇:PyTorch 自动微分是什么?理解 requires_grad 和 backward
  • 能区分输入、输出、数据和模型目标。

概念回顾

  • 【DataLoader】DataLoader 负责把数据一批一批喂给模型,训练循环才不会手忙脚乱。 前面或后面会反复用到它。
  • 【Linear 层】Linear 层做的是矩阵乘法加偏置,是很多神经网络模块的基础零件。 前面或后面会反复用到它。
  • 【Tensor】可以先把 Tensor 理解成支持 GPU、梯度和批量计算的多维数组。 前面或后面会反复用到它。

常见误区

  • 不要只记定义,要追问它解决了什么问题。
  • 不要只看 API 名字,要同时关注输入输出 shape 和训练流程。

课后练习

  • 用 3 句话向一个零基础朋友解释「计算图与梯度累积:为什么训练前必须清空梯度」。
  • 打开概念库里的「DataLoader」,补一遍它和本文的关系。
  • 读下一课「Shape 是深度学习最重要的数据契约」前,先写下你认为它会解决的问题。

自检练习与参考答案

1. 如果一段 PyTorch 代码报 device 不一致,第一步应该检查什么?

先检查模型参数、输入 Tensor、标签 Tensor 是否在同一个设备上。常见修正是把它们统一 `.to(device)`。

2. 为什么学习 PyTorch 时不能只看 API 名字?

因为训练是否正确经常取决于 shape、dtype、device 和梯度流,API 名字只能告诉你工具,不会保证数据契约正确。

3. 读完本文后,至少应该能画出哪条训练主线?

`输入数据 → forward → loss → backward → optimizer.step()`,并知道本文主题位于这条链路的哪一环。

下一步建议读「Shape 是深度学习最重要的数据契约」。

你大概已经写过上百遍这三行:

optimizer.zero_grad()
loss.backward()
optimizer.step()

但如果我问你:为什么第一行非写不可?去掉会怎样?很多人会说“梯度会累积”——可梯度为什么是累积的,而不是覆盖的?这背后是 PyTorch 计算图的设计哲学。搞不懂它,你遇到“loss 突然变 NaN”“训练前几步还好后面飞了”这类 bug,就只能瞎猜。

这一篇我们就把计算图、叶子节点和梯度累积彻底讲透。

第 45 课视频 - 计算图与梯度累积:为什么训练前必须清空梯度

概念回顾:上一节我们讲了什么

第 44 篇我们讲了 requires_gradbackward():当 requires_grad=True 时,PyTorch 会为这个 Tensor 的每个操作建计算图,backward() 沿着图反推梯度。今天我们往深走一步——这张图长什么样?backward() 之后它去哪了?梯度存在哪、怎么累加的?


一句话解释核心概念

PyTorch 用动态计算图记录前向运算,backward() 沿图反向算梯度并累加到参数的 .grad 上,算完图就销毁——所以下一个 batch 要先 zero_grad() 清掉旧梯度,否则梯度会像滚雪球一样越滚越大。


动态计算图到底是什么

你在第 14 篇学过前向传播和反向传播。前向传播是“算结果”,反向传播是“沿原路算梯度”。问题是:反向传播怎么知道“原路”长什么样?

答案就是计算图——它把前向传播的每一步运算都记下来。

来看代码:

import torch

# w 是模型参数(叶子节点),requires_grad=True
w = torch.tensor([2.0, 3.0], requires_grad=True)
# x 是输入数据,不需要梯度
x = torch.tensor([1.0, 2.0])

# 前向传播:每一步都被记录
y = (w * x).sum()   # y = 2*1 + 3*2 = 8

print(y.grad_fn)    # <SumBackward0 object at 0x...>

y.grad_fn 不是 None——这说明 PyTorch 给 y 挂了一个“反向传播函数”。这就是计算图在起作用。

关键术语

术语 含义 怎么认出它
叶子节点(leaf) 你直接创建的、不是算出来的 Tensor x.is_leaf == True
中间节点 由叶子节点运算产生的 Tensor grad_fn 不为 None
grad_fn 记录“这个 Tensor 是怎么算出来的” 反向传播的路线图
动态图 每次前向传播现建图,backward() 后销毁 和 TensorFlow 1.x 的静态图对比
w = torch.tensor([2.0], requires_grad=True)
print(w.is_leaf)        # True——直接创建,是叶子
print(w.grad_fn)        # None——叶子没有 grad_fn

y = w * 2
print(y.is_leaf)        # False——算出来的,不是叶子
print(y.grad_fn)        # <MulBackward0>

为什么强调“叶子节点”? 因为梯度只存在叶子节点的 .grad 里。中间节点的梯度用完就丢,不保留。


梯度累积:最坑也最重要的设计

现在来到这一篇的核心问题:为什么梯度是累加的?

PyTorch 的设计是:backward() 算出的梯度,不是覆盖到 .grad,而是加到 .grad 上。

w = torch.tensor([1.0], requires_grad=True)

# 第一次 backward
y1 = (w * 3).sum()
y1.backward()
print(w.grad)   # tensor([3.])  —— ∂y1/∂w = 3

# 第二次 backward,不清零
y2 = (w * 5).sum()
y2.backward()
print(w.grad)   # tensor([8.])  —— 3 + 5 = 8,累加了!

看到没?第二次的梯度 5 被加到了第一次的 3 上面,结果是 8,而不是覆盖成 5

为什么 PyTorch 要这么设计

这其实不是 bug,是 feature。有些场景需要梯度累积:

  • 显存不够,用小 batch 模拟大 batch:做 4 次 batch_size=16 的前向反向,梯度累加 4 次,等价于 batch_size=64 的一次更新。
  • 梯度累加训练:在显存受限时模拟更大有效 batch size。

但日常训练里,每个 batch 应该独立计算梯度、独立更新参数。如果不清零,上一个 batch 的梯度会“污染”当前 batch,导致:

实际梯度 = 当前batch梯度 + 之前所有batch的梯度

这会让更新方向完全错乱,loss 飙升或变 NaN。

正确写法

for x, y in train_loader:
    optimizer.zero_grad()     # ① 清空上一轮梯度
    logits = model(x)         # ② 前向传播(建图)
    loss = criterion(logits, y)
    loss.backward()           # ③ 反向传播(算梯度,图销毁)
    optimizer.step()          # ④ 用梯度更新参数

顺序记住:清 → 前 → 算 → 更


计算图的生命周期:用完即弃

PyTorch 的图是动态的——每次前向传播现搭,backward() 一调用就拆。

w = torch.tensor([1.0], requires_grad=True)
y = (w * 2).sum()

print(y.requires_grad)   # True
y.backward()

# backward 之后,图已销毁
# 再调用 backward 会报错
y.backward()   # RuntimeError: Trying to backward through the graph a second time

想多次 backward?得在前向时设 retain_graph=True

y = (w * 2).sum()
y.backward(retain_graph=True)   # 保留图
y.backward()                     # 第二次,这次之后图销毁

但日常训练不需要这个——每个 batch 都是全新的前向传播,建一张新图。


三个高频错误

错误 1:忘记 zero_grad

# ❌ 不报错,但训练完全跑偏
for x, y in train_loader:
    logits = model(x)
    loss = criterion(logits, y)
    loss.backward()
    optimizer.step()
    # 梯度越累越大,几轮后 loss 爆炸

症状:前几个 epoch 还正常,突然 loss 变 NaN 或飞到天上。

错误 2:backward 之后还想用图

loss.backward()
loss.backward()   # ❌ RuntimeError,图已经没了

修复:每个 batch 只 backward 一次。如果要算梯度又不想销毁图,用 retain_graph=True

错误 3:评估时建图白费显存

# ❌ 评估时还在建计算图
model.eval()
for x, y in test_loader:
    logits = model(x)        # 建了图,但根本不会 backward
    loss = criterion(logits, y)
    # 没调 backward,但图占着显存

修复

model.eval()
with torch.no_grad():        # ✅ 不建图,省显存提速度
    for x, y in test_loader:
        logits = model(x)
        loss = criterion(logits, y)

评估三件套:model.eval() + torch.no_grad() + 不调 backward()


detach():从图里把数据“摘”出来

有时候你想拿到一个 Tensor 的值,但不想让它参与梯度计算。用 detach()

w = torch.tensor([1.0], requires_grad=True)
y = w * 2

# y 还在图里
y_detached = y.detach()
# y_detached 脱离了图,requires_grad=False

print(y_detached.requires_grad)   # False
print(y_detached.grad_fn)         # None

常见用途:打印 loss 时别让 loss 留在图里占显存:

loss = criterion(logits, y)
print(loss.item())   # ✅ .item() 返回 Python float,自动脱离
# 或者
print(loss.detach())  # ✅ 显式脱离

课后练习

练习 1:下面这段代码运行后 w.grad 是多少?先猜再跑。

w = torch.tensor([1.0], requires_grad=True)
for i in range(5):
    y = (w * (i + 1)).sum()
    y.backward()
print(w.grad)

练习 2:如果训练循环里把 optimizer.zero_grad() 移到 optimizer.step() 后面(而不是最前面),训练还能正常进行吗?为什么?

练习 3:下面代码哪里有问题?怎么改?

model.eval()
for x, y in test_loader:
    logits = model(x)
    preds = torch.argmax(logits, dim=1)
    acc = (preds == y).float().mean()
    acc.backward()
参考答案 / 自检思路

练习 1:每次 backward 累加,梯度依次是 1+2+3+4+5=15。结果 tensor([15.])。这就是梯度累积的直观体现——不清零就会一直加。

练习 2:能正常训练。关键在于“每个 batch 更新前梯度是干净的”。放在 step() 后面等于“为下一个 batch 清零”,效果一样。但放在最前面是更常见的写法,因为语义更清晰:先清空,再用。如果放在 step 后面,第一个 batch 开始前没有清零(初始 grad 是 None,PyTorch 第一次 backward 时视为 0,所以第一个 batch 没问题),后续 batch 都能正常清零。

练习 3:两个问题。① 评估时不应该调 backward(),这里不需要梯度,白建图白算梯度,浪费显存。② acc 是评估指标,不是 loss,反向传播它毫无意义。改成:

model.eval()
with torch.no_grad():
    for x, y in test_loader:
        logits = model(x)
        preds = torch.argmax(logits, dim=1)
        acc = (preds == y).float().mean()

核心要点小结

  • PyTorch 用动态计算图记录前向运算,backward() 后图自动销毁
  • 叶子节点才有 .grad,中间节点的梯度用完即弃
  • 梯度是累加的——不清零就会和上一轮的梯度叠加,导致训练崩溃
  • 训练四步:清 → 前 → 算 → 更(zero_grad → forward → backward → step)
  • 评估时用 torch.no_grad() 避免建图,detach() 把 Tensor 从图中摘出来

下一篇我们聊 Shape——深度学习里最容易报错、也最重要的“数据契约”。你以为写错了维度只是报错,其实 Shape 错了模型可能“不报错但学废”。

Discussion

留言讨论