首页 / PyTorch 入门教程 / 梯度管理:叶子张量、zero_grad 与 detach

PyTorch 入门教程

梯度管理:叶子张量、zero_grad 与 detach

本教程共 60 篇 · 第 7 篇 · 更新于 2026-08-17 · 约 2 分钟阅读

梯度叶子张量zero_graddetachno_gradautogradPyTorch

本节目标:搞清叶子张量(leaf tensor)和非叶子张量的区别,掌握梯度清零的时机,学会用 detach()no_grad() 在不需要求导的地方停下,避开原地操作的坑。

上一章我们知道了 backward() 能把梯度算出来。但实际训练时,问题往往不是”怎么算”,而是”怎么管”:梯度存在哪、什么时候清零、什么时候干脆别算了。这一章就是解决这些管理问题的。

叶子与非叶子:谁是”原住民”

先做个区分。你自己创建的张量,比如 torch.tensor(...) 出来的,叫叶子张量。运算算出来的中间结果,叫非叶子张量。用 is_leaf 一眼分辨:

x = torch.ones(2, 2, requires_grad=True)  # 叶子
y = x + 2                                  # 非叶子
print(x.is_leaf)  # True
print(y.is_leaf)  # False

这个区分有啥用?关系到一条关键规则:

Note

backward() 之后,只有 requires_grad=True叶子张量会把梯度存进 .grad。非叶子张量的梯度用完就扔,.gradNone

这么设计是为了省内存。中间节点成百上千,梯度每个都存,显存很快就爆了。真需要查看某个中间张量的梯度(比如调试时),提前调一下 retain_grad()

x = torch.ones(2, 2, requires_grad=True)
y = x + 2
y.retain_grad()
z = (y * y).sum()
z.backward()
print(y.grad)  # tensor([[6., 6.], [6., 6.]]) 现在能看到了

平时训练用不到这个,知道它存在就行。

zero_grad:为什么必须清零

上一章说过,梯度是累积的。来看个活生生的坑:

x = torch.tensor(2.0, requires_grad=True)
for i in range(2):
    loss = x ** 2
    loss.backward()
    print(x.grad)  # 第一轮 tensor(4.),第二轮 tensor(8.)!

第二轮本该是 4,结果是 8——上次的梯度还赖在 .grad 里。梯度累积本身是特性(有的训练技巧专门利用它),但训练循环里每轮都得清零。

标准做法是调用优化器的 zero_grad()。优化器(optimizer)后面章节会正式介绍,现在先看它的位置:

import torch

w = torch.tensor([2.0, 3.0], requires_grad=True)
lr = 0.01

# 假设这是训练循环的一轮
for step in range(3):
    loss = (w ** 2).sum()      # 1. 前向:算损失
    loss.backward()            # 2. 反向:算梯度
    with torch.no_grad():
        w -= lr * w.grad       # 3. 更新参数
    w.grad.zero_()             # 4. 清零梯度,迎接下一轮
    print(step, loss.item())

四步走:前向、反向、更新、清零。以后用了优化器,第 3、4 步会换成 optimizer.step()optimizer.zero_grad(),流程不变。

Tip

清零除了 w.grad.zero_(),也可以写 w.grad = None。后者直接扔掉旧的梯度对象,比原地清零稍省一点内存,效果一样。

detach:把结果从图里摘出来

有时候你只是想要一个中间张量的数值,不想让它继续背着计算图。detach() 返回一个新张量,数据和原来共享,但和计算图彻底断开:

x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
y = x ** 2
y_detached = y.detach()
print(y_detached.requires_grad)  # False
print(y_detached.grad_fn)        # None

两个典型场景:

  1. 把张量转成 NumPy 数组。带梯度的张量不能直接转,必须先 detach:
arr = y.detach().numpy()
  1. 记录日志。循环里记 loss 时写成 loss.item();如果非要存张量,就存 loss.detach()。否则整个计算图会被一起留在列表里,显存越积越多,OOM 就是这么来的。

no_grad:整段代码暂停求导

比 detach 更省事的是 torch.no_grad() 上下文,包住的一整段代码都不建图:

with torch.no_grad():
    y = x ** 2
    print(y.requires_grad)  # False

模型评估、测试集预测时,把整段推理包起来,速度和内存都能省下一截。做推理时还有更强的 torch.inference_mode(),效果类似且更快,后面训练章节会再提。

两个容易踩的坑

第一,整数张量不支持求导。requires_grad=True 只对浮点类型有效,torch.tensor(2, requires_grad=True) 会直接报错,记得写成 2.0

第二,别对叶子张量做原地(in-place)修改。x += 1 这种写法可能破坏计算图,报错信息还很晦涩。要改就写 y = x + 1,新建一个张量。唯一安全的原地操作场合,是上面训练循环里 torch.no_grad() 包着的参数更新。

到这里,autograd 这套机制你就完整掌握了:会求梯度,会管梯度,也知道什么时候不求。下一章换个话题——设备管理,让代码从 CPU 跑到 GPU 上。