梯度管理:叶子张量、zero_grad 与 detach
本教程共 60 篇 · 第 7 篇 · 更新于 2026-08-17 · 约 2 分钟阅读
本节目标:搞清叶子张量(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。非叶子张量的梯度用完就扔,.grad是None。
这么设计是为了省内存。中间节点成百上千,梯度每个都存,显存很快就爆了。真需要查看某个中间张量的梯度(比如调试时),提前调一下 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
两个典型场景:
- 把张量转成 NumPy 数组。带梯度的张量不能直接转,必须先 detach:
arr = y.detach().numpy()
- 记录日志。循环里记 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 上。