🏠 总目录📚 本教程 04 · detach 与 no_grad ← →
📑 本页目录(点开跳转)

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)

实测输出:

关键信息

对非叶子设 requires_grad -> you can only change requires_grad flags of leaf variables. If you want to use a computed variable in a subgraph that doesn't require differentiation use var_no_grad = var.detach().
False op False -> False
False op True -> True
detach 出来的 : is_leaf = True 重新打开后 requires_grad = True

⭐ 这条报错原文自己把答案写出来了(这正是 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

⭐ 三行分别在说三件事:

  1. w 自己一点没变 —— 进了块它还是 requires_grad=True。no_grad 从来不改张量的属性。
  2. 变的是「新算出来的东西」 —— 块内的乘法不建图,所以 t1 身上没有回去的路。
  3. ⭐ 出了块立刻恢复 —— 同一行代码 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)

实测输出:

关键信息

建图 5 层后 v.grad_fn -> MulBackward0
no_grad 里 5 层后 -> None

⭐ 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 = [99.0, 2.0]
a.data_ptr() == b.data_ptr() -> True

⚠️⚠️ 改 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)

实测输出:

关键信息

detach 之后 backward -> element 0 of tensors does not require grad and does not have a grad_fn

⭐ 这条报错要认得 —— 它的意思是「你让我从一个没有图的东西开始反向」。 它有两个常见来源,⚠️ 报错原文完全一样,但病因不同:

来源 现场 修法
⭐ 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 或 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())

实测输出:

关键信息

inference_mode 里算的: requires_grad = False is_inference = True
拿去建图 -> Inference tensors cannot be saved for backward. Please do not use Tensors created in inference mode in computation tracked by autograd. To work around this, you can make a clone to get a normal tensor and use it in autograd, or use `torch.no_grad()` instead of `torch.inference_mode()`.
换成 no_grad -> 正常跑通, w2.grad = [6.0]

⭐ 区别就在这条报错里: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」是本章第六节的反面教材

✅ 检查点

  1. 用一句话分别说清 requires_grad / no_grad / detach 各自「断」的是什么?三者作用的对象分别是什么?
  2. 实测里 y.detach() 和 no_grad 里算的 x*3,三个字段完全一样。那它们的区别在哪?
  3. 为什么不能对一个中间结果设 requires_grad?那条报错原文自己给了什么建议?
  4. with torch.no_grad(): 块里,w.requires_grad 变了吗?块内算的和块外算的分别是什么?
  5. ⚠️ 在 no_grad 块里算出来的张量,拿到块外继续算,还能建图吗?为什么说「出了块就恢复」说的是代码不是数据?
  6. 💀 b = a.detach(); b[0] = 99 之后 a 是什么?要一份能随便改的副本该怎么写?
  7. ⭐ 想冻住网络中间的一层,A(requires_grad_(False))/ B(detach 输入)/ C(no_grad)三种写法,上游、中间、下游的梯度分别是什么?为什么只有一种是对的?
  8. 为什么 running_loss += loss.item() 里那个 .item() 不能省?实测 10 次循环不 detach 会挂多少个节点?
  9. inference_mode 和 no_grad 差在哪两点?为什么不建议在训练脚本里用前者?
  10. ⚠️ 冻主干时,requires_grad_(False) 漏掉了什么?只有哪个开关能管住它?
👀 答案
  1. requires_grad 标记的是这个叶子要不要梯度(对象是一个属性);no_grad 关的是当前这段代码建不建图(对象是一段代码);detach 剪的是这一个张量往回的那条边(对象是一个已经存在的张量)。三件事都被叫「断梯度」,尺度差了三个量级。
  2. 实测两行都是 requires_grad=False grad_fn=None is_leaf=True。区别在那次乘法建没建图:y = x * 3 建了(y 自己有 MulBackward0),detach 只是给它做了个不带图的替身;而 no_grad 里那次乘法根本没建图。一个作用在已经存在的张量上,一个作用在还没执行的代码上。
  3. 因为 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().
  4. w 一点没变,块内它仍然 requires_grad = True。变的是新算出来的东西:块内的 t1.requires_grad = False,块外的 t2.requires_grad = True —— 它是个作用域开关,不是标记。
  5. ❌ 不能。实测 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,不报错,但那段的参数一步都不动。
  6. 💀 a 变成 [99.0, 2.0] —— 实测 a.data_ptr() == b.data_ptr() 是 True,detach 不拷贝数据,只换说明书。要能随便改的副本必须写 .detach().clone()(detach 断图、clone 断内存,两个都要)。
  7. 实测:A 上游=有梯度、中间=None、下游=有梯度 只有它对;B 上游=None、中间=有梯度、下游=有梯度 💀 完全反了(detach 断的是输入那条边,层自己的参数是另一条边上的叶子);C 上游=None、中间=None、下游=有梯度 ⚠️ 连坐(那段没建图,反向走到就断,上游一起没了)。A 对的原因:它只关掉「在这层参数上累积梯度」,梯度照样穿过这一层流回上游。(若主干就在最前面、前面没有要训的东西,C 也对,而且额外省下那份反向存量。)
  8. 因为那是唯一切断这条链的地方。实测 10 次循环直接 total = total + loss:结果仍是 Tensor、requires_grad = True、图里 31 个节点(每轮 3 个 × 10 + 一个 AccumulateGrad),一个都放不掉;换成 .detach() 是 0 个节点,.item() 直接变成 Python float(值 55.0)。💀 真实训练里这表现为「跑到第几百个 batch 才 OOM」。⚠️ .item() 只能用在标量上,且会强制一次 GPU→CPU 同步。
  9. 两点:① 版本计数器 —— no_grad 照常维护,inference_mode 不维护(所以更快);② 产物能不能进后面的图 —— no_grad 的能(当常数,实测 w2.grad = [6.0]),inference_mode 的直接报 Inference tensors cannot be saved for backward。⚠️ 训练脚本里别用前者:一旦后来想拿验证时算的东西参与训练(自蒸馏、正则项)就会撞上,而且离现场很远。报错自己给了修法:clone() 或换回 no_grad()。
  10. ⚠️ 漏掉了 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

打卡记录保存在你的浏览器里,首页能看到总进度