📑 本页目录(点开跳转)
04 · detach / no_grad / requires_grad:三个都在断,断的不是一个东西
⏱ 94 分钟 | ⭐ 实测:想冻住中间那一层,三种写法后果全不一样 —— 一种把上游一起冻了,一种冻错了对象
🎯 一句话
requires_grad 标记的是「这个叶子要不要梯度」,no_grad 关的是「当前这段代码建不建图」,detach 剪的是「这一个张量往回的那条边」。
三件事都被人叫做「断梯度」,但断的对象分别是一个属性、一段代码、一条边 —— 尺度差了三个量级。
⭐ 本章存在的理由:ML 基础 附录A 里这三个词各有一行词条 ——
「no_grad():不建计算图,验证时用,省显存」「detach():切断梯度」「model.eval() vs no_grad():验证时两个都要」。
⚠️ 站内有这几行词条,但没有一页说清它们断的不是同一个东西 —— 于是「想冻住主干该用哪个」这种问题只能靠猜,
而第五节的实测会告诉你:三个写法里有两个会给你一个你不想要的结果,并且都不报错。
(本章也接着 03 章第八节:那里实测过 no_grad 和 detach 都拦不住版本计数器。
那是从「它们拦不住什么」的角度说的,本章反过来讲它们各自拦得住什么。)
🧩 一、先把三个摆在一起看
import torch
def show(tag, t):
fn = type(t.grad_fn).__name__ if t.grad_fn is not None else "None"
print(f"{tag:14s} requires_grad={str(t.requires_grad):6s}"
f" grad_fn={fn:14s} is_leaf={t.is_leaf}")
x = torch.tensor([1.0, 2.0], requires_grad=True)
y = x * 3
show("y = x * 3", y)
show("y.detach()", y.detach())
with torch.no_grad():
show("no_grad 里 x*3", x * 3)
实测输出:
对照
y = x * 3 requires_grad=True grad_fn=MulBackward0 is_leaf=False
y.detach() requires_grad=False grad_fn=None is_leaf=True
no_grad 里 x*3 requires_grad=False grad_fn=None is_leaf=True
⭐ 后两行的三个字段一模一样 —— 这就是「它们看起来是同一件事」的来源。 但它们是两条完全不同的路走到同一个结果的:
y.detach() |
no_grad 里的 x * 3 |
|
|---|---|---|
| 那次乘法发生了吗 | ⭐ 发生了,而且建了图(y 自己有 grad_fn) |
⭐ 发生了,但没建图 |
| 它作用在谁身上 | ⭐ 一个已经存在的张量 | ⭐ 一段还没执行的代码 |
| 结果张量 | 是 y 的一个「不带图的替身」 |
是一个凭空生出来的叶子 |
| 什么时候用得上 | 你手里已经有个张量,想让它从这里往回不再回传 | 你接下来这一段根本不打算求导 |
⚠️ is_leaf=True 这一列容易误读:它不代表「这是参数」。
02 章第五节说过,叶子的定义是「不是由要求导的运算算出来的」——
detach 和 no_grad 的产物都符合这个定义,所以它们都是叶子,但它们的 requires_grad 是 False,也就拿不到 .grad。
🧩 二、requires_grad:一个属性,而且只能设在叶子上
import torch
x = torch.tensor([1.0], requires_grad=True)
y = x * 2
try:
y.requires_grad_(False)
except RuntimeError as e:
print("对非叶子设 requires_grad ->", e)
a = torch.tensor([1.0]) # 默认 requires_grad=False
b = torch.tensor([1.0], requires_grad=True)
print("False op False ->", (a * a).requires_grad)
print("False op True ->", (a * b).requires_grad)
d = (x * 2).detach()
d.requires_grad_(True) # ⭐ detach 出来的是叶子,能重新打开
print("detach 出来的 : is_leaf =", d.is_leaf, " 重新打开后 requires_grad =", d.requires_grad)
实测输出:
关键信息
⭐ 这条报错原文自己把答案写出来了(这正是 00 章立的主题 ④):
「想让一个算出来的变量进入一张不求导的子图,就用 var.detach()」——
PyTorch 官方在这里明确告诉你:requires_grad 是给叶子用的,中间结果请用 detach。
⭐ 三条规则:
| 规则 | 说明 |
|---|---|
| 只能设在叶子上 | 参数、你自己 torch.tensor(..., requires_grad=True) 造的那些。中间结果一律拒绝 |
| ⭐ 中间结果的它是【算出来的】 | 判据是 OR:只要任何一个输入要梯度,输出就要(实测 False op True -> True) |
| 它是个持久属性 | 设一次就一直是那样,⚠️ 不像 no_grad 出了块就恢复(第三节实测) |
⚠️ requires_grad=False 的叶子会怎样:它照样参与前向计算、照样出现在表达式里,
只是图里不给它挂 AccumulateGrad 节点(02 章第一节),
反向走到它就停 —— 所以它的 .grad 永远是 None。
🧩 三、no_grad:它是个开关,不是个标记
⭐ 最常见的误解:以为 with torch.no_grad(): 会把块里碰到的张量「设成不要梯度」。它不会。
import torch
w = torch.tensor([1.0], requires_grad=True)
with torch.no_grad():
print("块内 w.requires_grad =", w.requires_grad)
t1 = w * 2
print("块内算的 t1.requires_grad =", t1.requires_grad)
t2 = w * 2
print("块外算的 t2.requires_grad =", t2.requires_grad)
实测输出:
要点
块内 w.requires_grad = True
块内算的 t1.requires_grad = False
块外算的 t2.requires_grad = True
⭐ 三行分别在说三件事:
w自己一点没变 —— 进了块它还是requires_grad=True。no_grad从来不改张量的属性。- 变的是「新算出来的东西」 —— 块内的乘法不建图,所以
t1身上没有回去的路。 - ⭐ 出了块立刻恢复 —— 同一行代码
w * 2写在块外就正常建图。它是个作用域开关,不是标记。
它到底省了什么
import torch
x = torch.tensor([1.0], requires_grad=True)
v = x
for _ in range(5):
v = v * 2
print("建图 5 层后 v.grad_fn ->", type(v.grad_fn).__name__)
with torch.no_grad():
u = x
for _ in range(5):
u = u * 2
print("no_grad 里 5 层后 ->", u.grad_fn)
实测输出:
关键信息
⭐ None 的意思是「根本没建」,不是「建了但断开了」。
02 章第七节量过:一次前向为反向存下来的东西是 44.0 MB 量级。
no_grad 省的就是这一份 —— 这才是「验证时用 no_grad 省显存」的真正含义,
它省的不是激活值本身,是为了反向而额外留下的那一份引用。
⚠️ 一个容易踩空的地方:从块里带出来的张量
import torch
w = torch.tensor([2.0], requires_grad=True)
with torch.no_grad():
a = w * 3
print("no_grad 出来的 a: requires_grad =", a.requires_grad)
b = a * 2 # ⚠️ 出了块,但 a 身上没有图
print("a * 2 呢 : requires_grad =", b.requires_grad, " grad_fn =", b.grad_fn)
(a * w).sum().backward() # ⭐ a 当常数用,w 那条路照样能回传
print("a * w 回传后 : w.grad =", w.grad.tolist())
实测输出:
对照
no_grad 出来的 a: requires_grad = False
a * 2 呢 : requires_grad = False grad_fn = None
a * w 回传后 : w.grad = [6.0]
⚠️ 「出了块就恢复」说的是【代码】,不是【数据】。 a 是在块里生出来的,它身上永远没有回去的路 ——
拿它出来再算 a * 2,结果照样 requires_grad=False。
⭐ 但这不影响别的路径:a * w 里 w 还是要梯度的,实测 w.grad = [6.0](a 的值是 6,被当成常数)。
💀 典型事故:在 no_grad 里算了个 embedding 或者一段特征,出来接着拿它算 loss,
loss 能算出来、backward() 也不报错,但那一段的参数一步都不会动。
🧩 四、detach:剪的是边,不是内存
⭐ detach() 不拷贝数据。它返回的是同一块内存的另一个「说明书」 ——
这是 01 章那条「一块 storage + 一份说明书」的直接后果。
import torch
a = torch.tensor([1.0, 2.0], requires_grad=True)
b = a.detach()
b[0] = 99.0
print("改了 b 之后 a =", a.tolist())
print("a.data_ptr() == b.data_ptr() ->", a.data_ptr() == b.data_ptr())
实测输出:
关键信息
⚠️⚠️ 改 b 把 a 改掉了,而 a 是个要梯度的参数。
💀 这是「我 detach 出来一份自己改改,应该不影响原件吧」的头号翻车现场 ——
它影响,而且是完全影响。
⭐ 要一份能随便改的副本,必须 .detach().clone()。(两个都要:detach 断图,clone 断内存。)
⚠️ 顺序上 detach().clone() 比 clone().detach() 略省一点(前者的 clone 不进图),实践中差别很小,但习惯写前者。
断掉之后就回不去了
import torch
x = torch.tensor([2.0], requires_grad=True)
y = (x * 3).detach()
try:
y.sum().backward()
except RuntimeError as e:
print("detach 之后 backward ->", e)
实测输出:
关键信息
⭐ 这条报错要认得 —— 它的意思是「你让我从一个没有图的东西开始反向」。 它有两个常见来源,⚠️ 报错原文完全一样,但病因不同:
| 来源 | 现场 | 修法 |
|---|---|---|
⭐ loss 上有 detach() / .item() |
中间某一步为了「省显存」把图断了 | 去掉那个 detach |
⭐ 整个前向在 no_grad 里 |
💀 最常见:验证循环的 no_grad 块写得太大,把训练那几行也包进去了 |
把 backward() 挪出块外,或缩小块的范围 |
⚠️ 还有第三种更隐蔽的:输入和参数全都 requires_grad=False —— 比如你把整个模型都冻住了,那 loss 本来就没有图。
🧩 五、⭐ 想冻住一层,到底该用哪个(最容易写反)
先看单独一层的情况:
import torch
import torch.nn as nn
torch.manual_seed(0)
inp = torch.randn(1, 4)
lin1 = nn.Linear(4, 4)
lin1.requires_grad_(False) # 写法 A:把这一层的参数标成不要梯度
print("A lin.requires_grad_(False) out.requires_grad =", lin1(inp).requires_grad)
lin2 = nn.Linear(4, 4)
out2 = lin2(inp.detach()) # 写法 B:把【输入】断掉
print("B lin(inp.detach()) out.requires_grad =", out2.requires_grad)
lin3 = nn.Linear(4, 4)
with torch.no_grad(): # 写法 C:整段不建图
out3 = lin3(inp)
print("C with no_grad(): lin(inp) out.requires_grad =", out3.requires_grad)
实测输出:
对照
A lin.requires_grad_(False) out.requires_grad = False
B lin(inp.detach()) out.requires_grad = True
C with no_grad(): lin(inp) out.requires_grad = False
⭐ B 是那个「看起来最像断梯度、其实什么都没冻住」的写法:
inp.detach() 断的是输入那条路,⚠️ 这一层自己的 weight 和 bias 照样要梯度,
所以输出 requires_grad=True,反向一走照样更新它。你以为冻了主干,其实只是不让梯度往输入方向流。
放进一个三层网络里,差别才完全显出来
单层看不出 A 和 C 的区别(都是 False)。⭐ 把「想冻的那层」夹在中间,三者立刻分道扬镳:
import torch
import torch.nn as nn
def run(tag, how):
torch.manual_seed(0)
up = nn.Linear(4, 4) # 上游:要训
mid = nn.Linear(4, 4) # 中间:想冻住的那一层
head = nn.Linear(4, 1) # 下游:要训
h = up(torch.randn(1, 4))
if how == "A":
mid.requires_grad_(False)
m = mid(h)
elif how == "B":
m = mid(h.detach())
else:
with torch.no_grad():
m = mid(h)
head(m).sum().backward()
mark = lambda layer: "有梯度" if layer.weight.grad is not None else "None "
print(f"{tag:32s} 上游={mark(up)} 中间={mark(mid)} 下游={mark(head)}")
run("A mid.requires_grad_(False)", "A")
run("B mid(h.detach())", "B")
run("C with no_grad(): mid(h)", "C")
实测输出:
对照
A mid.requires_grad_(False) 上游=有梯度 中间=None 下游=有梯度
B mid(h.detach()) 上游=None 中间=有梯度 下游=有梯度
C with no_grad(): mid(h) 上游=None 中间=None 下游=有梯度
⭐ 一张表读完这三行:
| 写法 | 上游(要训) | 中间(想冻) | 下游(要训) | 判定 |
|---|---|---|---|---|
A mid.requires_grad_(False) |
⭐ 有梯度 | None | 有梯度 | ⭐ 只有它做对了「冻住中间」这件事 |
B mid(h.detach()) |
💀 None | ⚠️ 有梯度 | 有梯度 | 💀 完全反了 —— 该冻的没冻,不该冻的冻了 |
C with no_grad(): mid(h) |
💀 None | None | 有梯度 | ⚠️ 连坐:它把整条回上游的路一起掐了 |
⭐ 原因一句话:
- A 只关掉「在这一层的参数上累积梯度」这一件事,梯度照样穿过这一层流回上游 —— 这正是你要的。
- B 断的是输入那条边,所以梯度回不到上游;而
mid自己的参数是另一条边上的叶子,没被碰到。 - C 那一段代码根本没建图,反向走到
m就没路了 —— ⚠️ 上游的参数一起变成None。
📋 那么实践里该怎么写
| 场景 | 用哪个 | 为什么 |
|---|---|---|
| ⭐ 冻主干、只训新头(主干在最前面) | A 或 C 都对 | ⭐ 主干前面没有要训的东西,C 的「连坐」无害;⭐ C 还额外省了那份反向存量(第三节的 44.0 MB 那一份) |
| ⭐ 冻网络【中间】的一段(前后都要训) | ⭐ 只能用 A | C 会把上游一起掐掉,B 冻错对象 |
| 只想让梯度不往某个输入回传(不是冻参数) | B(detach 输入) |
⭐ 这才是 detach 的正确用途 —— 第六节的目标网络和累加 loss 都是这一类 |
⚠️ A 还有一件必须配套做的事:把参数交给优化器时别把冻住的那些也交进去。
ML 基础 17 章「方案 B:冻结主干」用的正是写法 A,
它的优化器只收了 m.fc.parameters() —— 这不是随手写的,是配套的一半。
⚠️⚠️ 写法 A 冻的只是「参数的梯度」,不是「这一层的全部行为」 —— 第八节讲它漏了什么。
🛑 读到这里可以停 —— 前半章讲完了(约 42 分钟)。 后半章还有:日常最常撞的两处 ·
inference_mode:比no_grad更狠的那个 · ⚠️ 这三个都断不了的那件事 · 反查表 回来的时候不用重读,直接从下一节接着看就行。
🧩 六、日常最常撞的两处
① 累加 loss 做统计:不 detach 就是把整个 epoch 的图挂在一根线上
import torch
def count_nodes(t):
"""数一数从 t 往回能走到多少个反向节点。"""
if not torch.is_tensor(t) or t.grad_fn is None:
return 0
seen, stack = set(), [t.grad_fn]
while stack:
fn = stack.pop()
if fn is None or fn in seen:
continue
seen.add(fn)
for nxt, _ in fn.next_functions:
stack.append(nxt)
return len(seen)
x = torch.tensor([1.0], requires_grad=True)
total = 0.0
for i in range(10):
total = total + (x * (i + 1)).sum() # 💀 图一直挂着
print("直接累加 :", type(total).__name__, " requires_grad =", total.requires_grad,
" 图里节点数 =", count_nodes(total))
total2 = 0.0
for i in range(10):
total2 = total2 + (x * (i + 1)).sum().detach() # ⭐ 剪掉那条边
print("累加 detach :", type(total2).__name__, " requires_grad =", total2.requires_grad,
" 图里节点数 =", count_nodes(total2))
total3 = 0.0
for i in range(10):
total3 = total3 + (x * (i + 1)).sum().item() # ⭐ 直接变成 Python float
print("累加 .item() :", type(total3).__name__, " 值 =", total3)
实测输出:
对照
直接累加 : Tensor requires_grad = True 图里节点数 = 31
累加 detach : Tensor requires_grad = False 图里节点数 = 0
累加 .item() : float 值 = 55.0
💀💀 10 次循环,图里挂了 31 个节点,而且一个都放不掉。
每一轮 3 个(乘、求和、相加)×10 再加一个 AccumulateGrad。
⭐ 换成真实训练:这里的每一轮是一个 batch,那 31 就变成「每个 batch 的整张图 × 一个 epoch 的 batch 数」,
每一张都拖着 02 章量的那 44.0 MB 量级的存量。
⭐ 这就是 running_loss += loss.item() 里那个 .item() 的全部意义 ——
它不是「为了打印好看」,是唯一切断这条链的地方。
| 写法 | 结果类型 | 图里节点 | 判定 |
|---|---|---|---|
total += loss |
Tensor | 💀 31(越跑越多) | ❌ 显存持续涨,典型表现是「跑到第几百个 batch 才 OOM」 |
total += loss.detach() |
Tensor | ⭐ 0 | ✅ 可以,还留在 GPU 上不用同步 |
⭐ total += loss.item() |
⭐ float | — | ✅ 最常见的写法。⚠️ 它会强制一次 GPU→CPU 同步,循环里频繁调有开销 |
⚠️ .item() 只能用在标量上(一个元素)。要留整个张量做统计用 .detach()。
② 目标不该被训:no_grad 包住目标那一段
强化学习基础 08 章的 DQN 里,TD 目标是这么算的:
import torch
q_pred = torch.tensor([1.0], requires_grad=True) # 当前网络的输出(要训)
r, gamma, q_next = 0.5, 0.99, torch.tensor([2.0]) # 目标网络的输出
with torch.no_grad(): # ⭐ 目标那一段不建图
tgt = r + gamma * q_next
loss = ((q_pred - tgt) ** 2).sum()
loss.backward()
print("目标是常数时 q_pred.grad =", q_pred.grad.tolist(), " tgt.requires_grad =", tgt.requires_grad)
实测输出:
要点
目标是常数时 q_pred.grad = [-2.9600000381469727] tgt.requires_grad = False
⭐ 这里 no_grad 和 detach 是等价的(写 tgt = (r + gamma * q_next).detach() 效果一样),
挑哪个看你顺手 —— ⭐ 一段用 no_grad,一个张量用 detach。
⚠️ 那一章的坑表里第一条就是「目标里忘了 torch.no_grad() → 梯度回传到目标网络,训练直接乱掉」。
⭐ 同一个模式在别处还有两个化身:GAN 训判别器时对生成器的输出 detach();
知识蒸馏时教师模型的 logits detach()。判据都是一句话:这个东西是【标签】,不是【预测】。
🧩 七、inference_mode:比 no_grad 更狠的那个
import torch
w = torch.tensor([2.0], requires_grad=True)
with torch.inference_mode():
a = w * 3
print("inference_mode 里算的: requires_grad =", a.requires_grad,
" is_inference =", a.is_inference())
try:
(a * w).sum().backward() # 💀 拿它去参与一张要求导的图
except RuntimeError as e:
print("拿去建图 ->", e)
w2 = torch.tensor([2.0], requires_grad=True)
with torch.no_grad():
a2 = w2 * 3
(a2 * w2).sum().backward() # ⭐ 换成 no_grad 就没这个限制
print("换成 no_grad -> 正常跑通, w2.grad =", w2.grad.tolist())
实测输出:
关键信息
⭐ 区别就在这条报错里:no_grad 出来的张量还能被后面的图当常数用(实测 w2.grad = [6.0]),
inference_mode 出来的张量带一个 is_inference 标记,一旦进入要求导的计算就直接拒绝。
它换来的是连版本计数器都不维护(03 章第二节那个计数器),所以更快一点。
torch.no_grad() |
torch.inference_mode() |
|
|---|---|---|
| 建图 | 不建 | 不建 |
| 版本计数器 | ⭐ 照常维护 | ⭐ 不维护(所以更快) |
| 产物能不能进后面的图 | ⭐ 能(当常数) | 💀 不能,直接报错 |
| 什么时候用 | ⭐ 默认选它,尤其训练脚本里的验证循环 | 纯推理服务:进来一个请求、出去一个结果,产物不会再回到训练里 |
⚠️ 别为了那点速度在训练脚本里用 inference_mode —— 一旦你后来想拿验证时算的某个东西参与训练
(比如做个自蒸馏、算个正则项),就会撞上这条报错,而且它离现场很远。
⭐ 报错原文自己给了两条修法:clone() 一份,或者换回 no_grad()。
🧩 八、⚠️ 这三个都断不了的那件事
⭐ 它们全都只管 autograd。凡是不走 autograd 的状态更新,一个都拦不住。
最典型的就是 BatchNorm 的 running_mean / running_var:它们是在前向里用滑动平均直接改写的,
既不是梯度、也不经过反向 —— 所以:
| 你以为能拦住的 | 实际 |
|---|---|
with torch.no_grad(): |
❌ 拦不住 —— 它只管建不建图,前向照跑 |
bn.requires_grad_(False) |
❌ 拦不住 —— 它只管参数要不要梯度,running_* 根本不是参数(是 buffer) |
x.detach() 喂进去 |
❌ 拦不住 —— 输入断没断图,跟统计量更新无关 |
⭐ bn.eval() |
✅ 只有它能 —— 它切的是那个 training 布尔 |
💀 这就是 ML 基础 17 章那条坑「冻结主干时忘了 bn.eval() → BN 的 running stats 被新数据带跑」的完整机制:
你按第五节写法 A 把主干 requires_grad_(False) 了,参数确实一个都没动,但 BN 的统计量在每次前向里被悄悄改掉了,
⚠️ 零报错、零梯度,模型的行为却在漂。
⭐ 07 章整章讲这件事,它的判据是:
eval() 管「层的行为」,这三个管「梯度」,两套开关正交,验证时两套都要开。
🛑 读到这里可以停 —— 已经读了约 64 分钟。 最后一段还有(约 26 分钟):反查表 · 检查点与走神救援 回来的时候不用重读,直接从下一节接着看就行。
📋 九、反查表
| 你想干的事 | 用哪个 | ⚠️ 别用哪个 |
|---|---|---|
| 验证 / 推理循环 | ⭐ with torch.no_grad(): + model.eval()(两个都要) |
只写一个 |
| 把一个张量当成「标签 / 常数」 | ⭐ t.detach() |
t.data(03 章第七节:静默算错) |
| 累加 loss 做统计 | ⭐ .item()(标量)/ .detach()(张量) |
直接 += loss(实测图涨到 31 个节点) |
| 冻住主干、只训新头 | ⭐ requires_grad_(False) + 优化器只收要训的参数 |
detach 输入(冻错对象) |
| 冻住网络中间一段 | ⭐ 只能 requires_grad_(False) |
no_grad(上游一起被掐) |
| 要一份能随便改的副本 | ⭐ .detach().clone() |
只写 .detach()(实测改副本会改原件) |
| 冻主干时连 BN 一起冻住 | ⭐ bn.eval() |
以为 requires_grad_(False) 够了 |
| 报错原文 | 病因 |
|---|---|
element 0 of tensors does not require grad and does not have a grad_fn |
从一个没有图的张量上 backward():loss 上有 detach,或 no_grad 块包大了 |
you can only change requires_grad flags of leaf variables |
对中间结果设 requires_grad —— ⭐ 报错原文自己说了:改用 detach() |
Inference tensors cannot be saved for backward |
inference_mode 的产物被拿去参与求导 —— ⭐ clone() 一份,或换回 no_grad() |
🔗 这一章连到哪里
| 相关的地方 | 为什么 |
|---|---|
| 01 · 张量到底是什么 | ⭐ 第四节「detach 改副本会改原件」的根源在那一章:一块 storage + 一份说明书,detach 换的只是说明书 |
| 02 · autograd 怎么建图 | ⭐ 本章第一节说的「叶子的定义」和第二节说的 AccumulateGrad,分别在它的第五节和第一节;第三、六节两次引的 44.0 MB 是它第七节量出来的 |
| 03 · 就地操作:三种报错 | ⭐ 它第八节实测 no_grad 和 detach 都拦不住版本计数器,本章讲它们各自拦得住什么;.data 为什么不能用也在它第七节 |
| 05 · nn.Module 内部 | 第五节写法 A 靠 requires_grad_(False) 一次走遍所有参数 —— parameters() 到底是怎么把它们找齐的,在那一章 |
| 07 · train / eval 改了什么 | ⭐ 第八节那件「三个都断不了的事」是它的正题:eval() 管层的行为,本章这三个管梯度,两套开关正交 |
| 09 · 自定义算子 | 第七节那条报错说的「不能被 saved for backward」—— 「什么叫 saved for backward、前向到底存了什么」在那一章 |
| ML 基础 附录A · 术语速查卡 | ⭐ 那里 no_grad() / detach() / eval() vs no_grad() 各一行词条,本章就是那三行的展开 |
| ML 基础 17 · 实战项目 | ⭐ 它「方案 B:冻结主干」用的正是本章第五节写法 A,优化器只收 fc 的参数;它那条「冻结时忘了 bn.eval()」的机制在本章第八节 |
| ML 基础 11 · 训练调试手册 | 「梯度断流四来源」第一条就是 .detach() / .data —— 本章第四节第二段那条报错就是它在现场的样子 |
| 强化学习基础 08 · DQN | ⭐ 第六节那个目标网络的实例出自那里;它的坑表第一条「目标里忘了 no_grad」是本章第六节的反面教材 |
✅ 检查点
- 用一句话分别说清
requires_grad/no_grad/detach各自「断」的是什么?三者作用的对象分别是什么? - 实测里
y.detach()和no_grad里算的x*3,三个字段完全一样。那它们的区别在哪? - 为什么不能对一个中间结果设
requires_grad?那条报错原文自己给了什么建议? with torch.no_grad():块里,w.requires_grad变了吗?块内算的和块外算的分别是什么?- ⚠️ 在
no_grad块里算出来的张量,拿到块外继续算,还能建图吗?为什么说「出了块就恢复」说的是代码不是数据? - 💀
b = a.detach(); b[0] = 99之后a是什么?要一份能随便改的副本该怎么写? - ⭐ 想冻住网络中间的一层,A(
requires_grad_(False))/ B(detach输入)/ C(no_grad)三种写法,上游、中间、下游的梯度分别是什么?为什么只有一种是对的? - 为什么
running_loss += loss.item()里那个.item()不能省?实测 10 次循环不 detach 会挂多少个节点? inference_mode和no_grad差在哪两点?为什么不建议在训练脚本里用前者?- ⚠️ 冻主干时,
requires_grad_(False)漏掉了什么?只有哪个开关能管住它?
👀 答案
requires_grad标记的是这个叶子要不要梯度(对象是一个属性);no_grad关的是当前这段代码建不建图(对象是一段代码);detach剪的是这一个张量往回的那条边(对象是一个已经存在的张量)。三件事都被叫「断梯度」,尺度差了三个量级。- 实测两行都是
requires_grad=False grad_fn=None is_leaf=True。区别在那次乘法建没建图:y = x * 3建了(y自己有MulBackward0),detach只是给它做了个不带图的替身;而no_grad里那次乘法根本没建图。一个作用在已经存在的张量上,一个作用在还没执行的代码上。 - 因为
requires_grad是给叶子用的 —— 中间结果的这个标志是算出来的(判据是 OR:任一输入要梯度,输出就要,实测False op True -> True)。报错原文自己写了修法:If you want to use a computed variable in a subgraph that doesn't require differentiation use var_no_grad = var.detach(). w一点没变,块内它仍然requires_grad = True。变的是新算出来的东西:块内的t1.requires_grad = False,块外的t2.requires_grad = True—— 它是个作用域开关,不是标记。- ❌ 不能。实测
a = w * 3(块内)出来后a.requires_grad = False,块外再算a * 2仍然是False、grad_fn = None——a身上永远没有回去的路。但别的路径不受影响:(a * w).backward()照样给出w.grad = [6.0](a当常数)。💀 典型事故:在no_grad里算了一段特征再拿去算 loss,不报错,但那段的参数一步都不动。 - 💀
a变成[99.0, 2.0]—— 实测a.data_ptr() == b.data_ptr()是True,detach不拷贝数据,只换说明书。要能随便改的副本必须写.detach().clone()(detach断图、clone断内存,两个都要)。 - 实测:A 上游=有梯度、中间=None、下游=有梯度 只有它对;B 上游=None、中间=有梯度、下游=有梯度 💀 完全反了(
detach断的是输入那条边,层自己的参数是另一条边上的叶子);C 上游=None、中间=None、下游=有梯度 ⚠️ 连坐(那段没建图,反向走到就断,上游一起没了)。A 对的原因:它只关掉「在这层参数上累积梯度」,梯度照样穿过这一层流回上游。(若主干就在最前面、前面没有要训的东西,C 也对,而且额外省下那份反向存量。) - 因为那是唯一切断这条链的地方。实测 10 次循环直接
total = total + loss:结果仍是Tensor、requires_grad = True、图里 31 个节点(每轮 3 个 × 10 + 一个AccumulateGrad),一个都放不掉;换成.detach()是 0 个节点,.item()直接变成 Pythonfloat(值 55.0)。💀 真实训练里这表现为「跑到第几百个 batch 才 OOM」。⚠️.item()只能用在标量上,且会强制一次 GPU→CPU 同步。 - 两点:① 版本计数器 ——
no_grad照常维护,inference_mode不维护(所以更快);② 产物能不能进后面的图 ——no_grad的能(当常数,实测w2.grad = [6.0]),inference_mode的直接报Inference tensors cannot be saved for backward。⚠️ 训练脚本里别用前者:一旦后来想拿验证时算的东西参与训练(自蒸馏、正则项)就会撞上,而且离现场很远。报错自己给了修法:clone()或换回no_grad()。 - ⚠️ 漏掉了 BatchNorm 的
running_mean/running_var—— 它们在前向里用滑动平均直接改写,不是梯度也不经反向,所以no_grad、requires_grad_(False)、detach输入三个全都拦不住(running_*根本不是参数,是 buffer)。只有bn.eval()能,它切的是那个training布尔。💀 这就是 ML 基础 17 章那条坑「冻结主干时忘了bn.eval()→ running stats 被新数据带跑」的机制:参数一个没动,模型行为却在漂,零报错。
🛑 可以停在这里
⚡ 走神救援
⭐ 三个词都叫「断梯度」,断的分别是一个属性、一段代码、一条边。
requires_grad标记「这个叶子要不要梯度」,⚠️ 只能设在叶子上——⭐ 而它的报错原文自己会让你改用detach()。
no_grad关的是「当前这段代码建不建图」,⭐ 它是作用域开关不是标记:块内参数的标志没变,变的是块内新算出来的东西。⚠️ 但「出了块恢复」说的是代码不是数据——💀 在no_grad里算特征再拿去算 loss,不报错,那段参数一步不动。
detach剪的是「这一个张量往回的那条边」。⚠️⚠️ 它不拷贝数据——改 detach 出来的张量会直接改到原张量,要能改的副本必须写.detach().clone()。⭐ 本章最容易写反的一节:冻住网络【中间】那一层。 只有
requires_grad_(False)是对的——它只关掉「在这层参数上累积梯度」,梯度照样穿过这一层;用detach会完全反过来(断的是输入那条边,层自己的参数在另一条边上);用no_grad会连坐,整条回上游的路一起掐掉。主干在最前面时no_grad也对且更省显存,冻中间只能用第一种。日常两处:累加 loss 不断图,💀 真实表现是「跑到几百个 batch 才 OOM」;目标网络/教师模型用
no_grad包住,⭐ 判据是「这个东西是标签不是预测」。
inference_mode比no_grad更狠也更快,但它的产物不能再进要求导的图。
下一节 👉 05-nnModule内部.md