📑 本页目录(点开跳转)
09 · 自己写一个 autograd 算子
⏱ 114 分钟 | ⭐ 实测:把 8 层 tanh 的中间结果一个都不存、反向时重算一遍,图里存下的字节从 8.0 MB 掉到 1.0 MB,梯度一位不差
🎯 一句话
autograd.Function 是你和框架签的一份合同:前向那一格随便你怎么算(可以出去用 numpy、可以完全不建图、可以做一个数学上根本不可导的操作),代价是【反向那一格必须你自己填】。
第 02 章讲的是框架自动建图;这一章讲的是手动往图里插一个节点 —— 插进去之后,它和内置算子在 autograd 眼里没有任何区别。
🧩 一、什么时候真的需要它
先说清楚什么时候不需要,因为这是最常见的误用。
⚠️ 「我想让它更快」不在清单里。 如果你的操作只是把已有算子组合起来(x.sigmoid() * w + b 这种),autograd 自己求出来的导数既正确又不慢,手写一个 Function 只会多一层 Python 调用、多一个会写错的地方。
✅ 真正需要它的是这三种,每一种都是框架自己算不出来:
① 真实的梯度是 0,或者根本不存在
import torch
x = torch.tensor([0.2, 0.7, 1.4], requires_grad=True)
torch.round(x).sum().backward()
print("torch.round 的梯度:", x.grad.tolist())
实跑输出:
要点
torch.round 的梯度: [0.0, 0.0, 0.0]
⭐ round 是阶梯函数,除了跳变点之外导数处处为 0 —— 框架没算错,它算得完全正确。
但正确的梯度是 0 意味着:这一步之前的所有参数都收不到任何信号,整个网络的前半截等于没训。
于是有了一个业界通用的作弊法,叫直通估计(Straight-Through Estimator):前向照样取整,反向假装自己是恒等函数。
import torch
class RoundSTE(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
return torch.round(x) # 前向:真的取整
@staticmethod
def backward(ctx, g):
return g # ⭐ 反向:假装自己是恒等函数,梯度原样放过去
x = torch.tensor([0.2, 0.7, 1.4], requires_grad=True)
y = RoundSTE.apply(x)
y.sum().backward()
print("STE 的前向:", y.tolist(), " 梯度:", x.grad.tolist())
实跑输出:
要点
STE 的前向: [0.0, 1.0, 1.0] 梯度: [1.0, 1.0, 1.0]
⭐ 前向是取整(0/1/1),梯度是恒等(全 1)—— 前向和反向【故意不匹配】。 这在数学上是错的,在工程上是量化感知训练(QAT)、二值网络、离散隐变量的标准做法。你在 《AI基础设施》18 · 量化 里看到的「QAT 是在训练时插入伪量化节点」,那个节点的反向就是这么写出来的。
② 前向绕出框架去算了
import numpy as np
import torch
x = torch.tensor([1.0, 2.0], requires_grad=True)
y = torch.from_numpy(np.sqrt(x.detach().numpy())) # 中间绕出去用 numpy 算
print("y.requires_grad =", y.requires_grad, " y.grad_fn =", y.grad_fn)
y.sum().backward()
实跑输出:
y.requires_grad = False y.grad_fn = None
RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn
⭐ 一旦离开框架,图就断在那里。 numpy 不知道自己在被求导,from_numpy 回来的是一个全新的叶子张量,和 x 之间没有任何关系。
Function 就是那个把断口重新焊上的东西 —— 第二节的第三个例子会把这一段补完整。
同样的情形还包括:调一段自己写的 C++ / CUDA 核、调一个只有 Python 接口的第三方库、调一个跑在别的进程里的模拟器。
③ ⭐ 你想拿「反向时重算」去换显存
这是第四节的正题,也是本章唯一一个「优化」用途 —— 注意它优化的不是速度,是显存,而且是用速度去换。
🔧 二、合同的四条条款
最小的一个 Function:
import torch
class Square(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return x * x
@staticmethod
def backward(ctx, grad_out):
(x,) = ctx.saved_tensors
return grad_out * 2 * x
x = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
y = Square.apply(x)
y.sum().backward()
print("自定义:", x.grad.tolist())
x2 = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
(x2 * x2).sum().backward()
print("内置 :", x2.grad.tolist())
x3 = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
print("grad_fn:", Square.apply(x3).grad_fn)
实跑输出:
要点
自定义: [2.0, 4.0, 6.0]
内置 : [2.0, 4.0, 6.0]
grad_fn: <torch.autograd.function.SquareBackward object at 0x000001554BAA0E50>
⭐ grad_fn 叫 SquareBackward —— 和第 02 章里见到的 MulBackward0、AddBackward0 是同一类东西。你的类名 + Backward 就是它在图里的名字,报错信息里也会出现它,所以类名起得可读一点,将来定位问题会省事。
四条条款:
| 条款 | 内容 | 违反会怎样 |
|---|---|---|
| ① | forward / backward 都是 @staticmethod,调用走 .apply() |
见下面那条最误导人的报错 |
| ② | forward 的第一个参数是 ctx,它是这一次调用的储物柜 |
—— |
| ③ | backward 收到的梯度,形状 = forward 输出的形状 |
形状对不上会在加梯度那一步炸 |
| ④ | backward 返回值的个数 = forward 输入参数的个数(ctx 不算),且一一对应 |
returned an incorrect number of gradients |
⚠️ 一条会把你带偏的报错
import torch
class Square(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return x * x
@staticmethod
def backward(ctx, g):
(x,) = ctx.saved_tensors
return g * 2 * x
x = torch.tensor([1.0, 2.0], requires_grad=True)
Square()(x) # ⚠️ 写成了「实例化再调用」
实跑输出(截去了后半段 URL):
DeprecationWarning: <class '__main__.Square'> should not be instantiated. Methods on autograd functions are all static, so you should invoke them on the class itself.
RuntimeError: Legacy autograd function with non-static forward method is deprecated. Please use new-style autograd function with static forward method.
💀 报错说你的 forward 不是 static,而上面这个 forward 明明写了 @staticmethod。
真正的原因是那个 Square() —— 你实例化了它。正确写法永远是 Square.apply(x),一个括号的差别。
⚠️ 这条报错会让人回去反复检查装饰器,那是白费的;看到 non-static forward method 先去找有没有 ()。
④ 返回几个梯度:数着输入参数来
import torch
class ScaleBad(torch.autograd.Function):
@staticmethod
def forward(ctx, x, k):
ctx.k = k
return x * k
@staticmethod
def backward(ctx, g):
return g * ctx.k # ⚠️ forward 收了 2 个参数,这里只返回 1 个
x = torch.tensor([1.0, 2.0], requires_grad=True)
ScaleBad.apply(x, 3.0).sum().backward()
实跑输出:
RuntimeError: function ScaleBadBackward returned an incorrect number of gradients (expected 2, got 1)
✅ 修法是给非张量参数返回 None:
@staticmethod
def backward(ctx, g):
return g * ctx.k, None # ⭐ k 是个 float,它没有梯度,但位置要占住
⭐ None 的意思是「这个入参我不给梯度」,不是「我忘了」。超参数、形状、字符串、布尔开关,全都返回 None。
🚦 needs_input_grad:别算没人要的梯度
import torch
class Clip(torch.autograd.Function):
"""把 x 夹到 [lo, hi],反向只让区间内的梯度过去。"""
@staticmethod
def forward(ctx, x, lo, hi):
ctx.save_for_backward(x)
ctx.lo, ctx.hi = lo, hi
return x.clamp(lo, hi)
@staticmethod
def backward(ctx, g):
(x,) = ctx.saved_tensors
print(" needs_input_grad =", ctx.needs_input_grad)
mask = (x >= ctx.lo) & (x <= ctx.hi)
gx = g * mask if ctx.needs_input_grad[0] else None # ⭐ 没人要就别算
return gx, None, None # lo / hi 不是张量
x = torch.tensor([-2.0, 0.5, 2.0], requires_grad=True)
y = Clip.apply(x, -1.0, 1.0)
print("前向:", y.tolist())
y.sum().backward()
print("梯度:", x.grad.tolist())
实跑输出:
要点
前向: [-1.0, 0.5, 1.0]
needs_input_grad = (True, False, False)
梯度: [0.0, 1.0, 0.0]
⭐ ctx.needs_input_grad 是一个和输入参数等长的布尔元组(这里是 (True, False, False):只有 x 要梯度,两个 float 不要)。
被夹住的 -2.0 和 2.0 梯度是 0,区间内的 0.5 是 1 —— 这就是 clamp 该有的行为,你亲手写了出来。
💡 有多个张量输入时这条能省很多:冻结了 backbone(见 04 · detach / no_grad / requires_grad)之后,needs_input_grad[0] 就是 False,那一大坨矩阵乘法可以直接跳过。
🧯 把断掉的图焊回去
回到第一节那个 numpy 的例子,现在能补完整了:
import numpy as np
import torch
class NpSqrt(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
out = torch.from_numpy(np.sqrt(x.detach().numpy())) # 出去用 numpy 算
ctx.save_for_backward(out)
return out
@staticmethod
def backward(ctx, g):
(out,) = ctx.saved_tensors
return g / (2 * out) # d/dx √x = 1/(2√x)
x = torch.tensor([1.0, 4.0, 9.0], requires_grad=True)
y = NpSqrt.apply(x)
print("前向:", y.tolist(), " grad_fn:", type(y.grad_fn).__name__)
y.sum().backward()
print("自定义梯度:", x.grad.tolist())
x2 = torch.tensor([1.0, 4.0, 9.0], requires_grad=True)
torch.sqrt(x2).sum().backward()
print("内置梯度 :", x2.grad.tolist())
实跑输出:
对照
前向: [1.0, 2.0, 3.0] grad_fn: NpSqrtBackward
自定义梯度: [0.5, 0.25, 0.1666666716337204]
内置梯度 : [0.5, 0.25, 0.1666666716337204]
⭐ 注意 grad_fn 从 None 变成了 NpSqrtBackward —— 中间那段 numpy 计算对 autograd 来说变成了一个不透明的黑盒,但图是连着的。
这就是「怎么把一个外部实现接进 PyTorch」的 Python 侧答案。
⚠️ C++ / CUDA 侧的答案不在这里 —— 一个真正的算子是怎么用 TORCH_LIBRARY 注册进去、怎么跟 dispatcher 打交道,在 《框架底下是 C++》05 · 算子怎么注册进 PyTorch。这一章只到 Python 这一层为止。
🧷 三、ctx.save_for_backward(x) 和 ctx.x = x 的差别
这两行看起来只是写法不同,实际上差着一整套版本检查。
import torch
class SaveViaCtx(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.x = x # ⚠️ 直接挂在 ctx 上
return x * x
@staticmethod
def backward(ctx, g):
return g * 2 * ctx.x
class SaveProper(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x) # ⭐ 正规写法
return x * x
@staticmethod
def backward(ctx, g):
(x,) = ctx.saved_tensors
return g * 2 * x
for name, F in (("ctx.x ", SaveViaCtx), ("save_for_backward", SaveProper)):
a = torch.tensor([1.0, 2.0, 3.0], requires_grad=True)
b = a * 1.0 # 造一个非叶子张量,好就地改
y = F.apply(b)
b.mul_(10) # 💀 前向存下来的东西被就地改了
try:
y.sum().backward()
print(f"{name}: 没报错,算出的梯度 = {a.grad.tolist()}")
except Exception as e:
print(f"{name}: {type(e).__name__}: {e}")
实跑输出:
要点
ctx.x : 没报错,算出的梯度 = [20.0, 40.0, 60.0]
查看完整报错:反向传播所需的变量被原地修改
save_for_backward: RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation: [torch.FloatTensor [3]], which is output 0 of Mul, is at version 1; expected version 0 instead. Hint: enable anomaly detection to find the operation that failed to compute its gradient, with torch.autograd.set_detect_anomaly(True, check_nan=False).
💀💀 正确答案是 [2.0, 4.0, 6.0](去掉 b.mul_(10) 那一行,两种写法都给出这个数,和 (b*b).sum().backward() 一致)。
ctx.x 那一版没有报任何错,安安静静地把梯度算大了 10 倍。
⭐ 为什么会差 10 倍:ctx.x 存的是一个引用。b.mul_(10) 之后 b 的内容变成了 [10, 20, 30],反向读到的 2 * ctx.x 就成了 2 × 10a,正好 10 倍。而前向输出 y 是在改之前算的 —— 前向用的是旧值,反向用的是新值,一次计算里混了两个版本。
⭐ save_for_backward 多做的那件事就是给张量记一个版本号:存的时候记下当前版本,反向取的时候比对,对不上就抛上面那条报错。
这正是 03 · 就地操作:三种报错,三个不同的原因 里第二种报错的机制 —— 现在你看到了它是从哪一行代码里发出来的。
| 存什么 | 用哪个 |
|---|---|
| 张量(输入、输出、任何要在反向里用的张量) | ⭐ ctx.save_for_backward(...),反向用 ctx.saved_tensors 取 |
| 非张量(float、int、shape、bool 开关、字符串) | ctx.k = k 直接挂,反向 ctx.k 读 |
⚠️ 不要为了「省事」把张量挂 ctx。你省掉的是两个下划线,换来的是一类不报错、只算错的 bug —— 而这类 bug 在训练里表现为「loss 下降得有点奇怪」,能耗掉你一个星期。
💡 只在需要时存:save_for_backward 会让那个张量在反向跑完之前一直活着,存进去 = 显存里多留一份。反向用不到就别存 —— 这一条直接引出下一节。
🛑 读到这里可以停 —— 前半章讲完了(约 44 分钟)。 后半章还有:把「重算」写进反向:8.0 MB → 1.0 MB ·
gradcheck:不要靠眼睛检查导数 · 二阶导:默认能,除非你明说不能 · 这一章不讲什么 回来的时候不用重读,直接从下一节接着看就行。
♻️ 四、把「重算」写进反向:8.0 MB → 1.0 MB
这是自定义算子最有工程价值的一个用法。
先看问题:前向每做一步,autograd 就要为反向留下一份中间结果。做 8 次 tanh,就留 8 份。
⭐ 可以直接量出来:torch.autograd.graph.saved_tensors_hooks 会在每一个张量被存进图的时候回调一次,把它们的字节数加起来,就是「这次前向为了反向留了多少」。
import torch
DEPTH = 8
N = 512
n_fwd = 0
def block(x):
"""一段没有参数的前向:连着做 DEPTH 次 tanh,每一步都会留下一个中间结果。"""
global n_fwd
n_fwd += 1
for _ in range(DEPTH):
x = torch.tanh(x)
return x
class Recompute(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x) # ⭐ 只存入口,中间一个都不存
return block(x) # forward 里本来就没开 grad,不会建图
@staticmethod
def backward(ctx, g):
(x,) = ctx.saved_tensors
with torch.enable_grad(): # ⭐ 反向里重新建一次图
xx = x.detach().requires_grad_(True)
y = block(xx)
return torch.autograd.grad(y, xx, g)[0]
torch.manual_seed(0)
X0 = torch.randn(N, N)
def measure(fn):
"""量一量:这次前向一共往图里存了多少字节。"""
seen, total = set(), [0]
def pack(t):
if id(t) not in seen:
seen.add(id(t))
total[0] += t.numel() * t.element_size()
return t
x = X0.clone().requires_grad_(True)
with torch.autograd.graph.saved_tensors_hooks(pack, lambda t: t):
y = fn(x)
y.sum().backward()
return total[0] / 1024 / 1024, x.grad
n_fwd = 0
mb1, g1 = measure(block)
f1 = n_fwd
n_fwd = 0
mb2, g2 = measure(Recompute.apply)
f2 = n_fwd
print(f"普通写法 : 存了 {mb1:.1f} MB block() 执行了 {f1} 次")
print(f"重算写法 : 存了 {mb2:.1f} MB block() 执行了 {f2} 次")
print(f"梯度是否一致: {torch.allclose(g1, g2)} 最大差 {(g1 - g2).abs().max().item():.2e}")
实跑输出:
对照
普通写法 : 存了 8.0 MB block() 执行了 1 次
重算写法 : 存了 1.0 MB block() 执行了 2 次
梯度是否一致: True 最大差 0.00e+00
⭐ 三个数字,把这笔交易写得清清楚楚:
| 数字 | 意思 |
|---|---|
| 8.0 MB → 1.0 MB | 512×512 的 float32 正好是 1 MB。普通写法存了 8 份(DEPTH=8,每步一份),重算写法只存了入口那 1 份 |
block() 从 1 次变 2 次 |
省下来的显存是用一次额外的前向换的。这就是「重算换显存」里的「换」 |
最大差 0.00e+00 |
⭐ 梯度逐位相同,不是「近似相等」。重算不是近似方法,它算的是同一件事 |
🧯 两个 grad mode 的坑
import torch
class Probe(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
print(" forward 里 is_grad_enabled() =", torch.is_grad_enabled())
ctx.save_for_backward(x)
return x * x
@staticmethod
def backward(ctx, g):
print(" backward 里 is_grad_enabled() =", torch.is_grad_enabled())
(x,) = ctx.saved_tensors
return g * 2 * x
x = torch.tensor([1.0, 2.0], requires_grad=True)
print("外面 is_grad_enabled() =", torch.is_grad_enabled())
Probe.apply(x).sum().backward()
实跑输出:
对照
外面 is_grad_enabled() = True
forward 里 is_grad_enabled() = False
backward 里 is_grad_enabled() = False
⭐ 两条都反直觉,各解决一个问题:
forward里 grad 已经是关的。 所以你在forward里做什么都不会建图 —— 不需要自己加with torch.no_grad()(加了也无害,只是多余)。这也解释了第二节那个 numpy 例子为什么不用担心图被污染。- ⚠️
backward里 grad 也是关的。 所以重算那一段必须用with torch.enable_grad():把它打开,否则block(xx)的结果没有grad_fn,torch.autograd.grad直接抛:
RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn
⚠️ 注意这条报错和第一节 numpy 那个例子是同一句话 —— 它的含义永远是「你要求导的那个东西不在图上」,但成因至少有三种:忘了 requires_grad=True、中途绕出了框架、或者这里的忘了 enable_grad。附录 A 把这条单列了一行。
💡 这个 Function,PyTorch 已经替你写好了
import torch
from torch.utils.checkpoint import checkpoint
# ... block / measure 同上 ...
mb, g = measure(lambda x: checkpoint(block, x, use_reentrant=False))
print(f"torch.utils.checkpoint: 存了 {mb:.1f} MB block() 执行了 {n_fwd} 次")
实跑输出:
要点
torch.utils.checkpoint: 存了 1.0 MB block() 执行了 2 次
⭐ 和手写的那个一模一样:1.0 MB、2 次。 实际项目里请直接用 torch.utils.checkpoint(它还处理了 RNG 状态、autocast 状态、多输出等一堆你不想自己写的细节)。
手写一遍的价值在于:你现在知道了 checkpoint 不是什么魔法,它就是一个「前向不存、反向重跑」的 Function,以及它为什么会让你的训练变慢。
👀 `torch.utils.checkpoint` 那段的完整可跑版(点开复制)
import torch
from torch.utils.checkpoint import checkpoint
DEPTH, N = 8, 512
n_fwd = 0
def block(x):
global n_fwd
n_fwd += 1
for _ in range(DEPTH):
x = torch.tanh(x)
return x
torch.manual_seed(0)
X0 = torch.randn(N, N)
def measure(fn):
seen, total = set(), [0]
def pack(t):
if id(t) not in seen:
seen.add(id(t))
total[0] += t.numel() * t.element_size()
return t
x = X0.clone().requires_grad_(True)
with torch.autograd.graph.saved_tensors_hooks(pack, lambda t: t):
y = fn(x)
y.sum().backward()
return total[0] / 1024 / 1024, x.grad
n_fwd = 0
mb, g = measure(lambda x: checkpoint(block, x, use_reentrant=False))
print(f"torch.utils.checkpoint: 存了 {mb:.1f} MB block() 执行了 {n_fwd} 次")
⚠️ 这一节只讲【怎么把重算写进反向】这个框架接口。 「激活到底占了显存的多少、该在哪里切、切多少段最划算、和 ZeRO/FSDP 怎么配」是算法与显存账,在 《AI基础设施》09 · 显存优化全家桶; FlashAttention 那种「连注意力矩阵都不存、反向现算」的极端形态在 《AI基础设施》08 · FlashAttention。⭐ 那两章回答「为什么值得」,这一章回答「怎么写」。
🛑 第二个休息点 —— 中段讲完了(约 22 分钟)。 最后一段还有:
gradcheck:不要靠眼睛检查导数 · 二阶导:默认能,除非你明说不能 · 这一章不讲什么 这一章确实长,分三次读完全没问题 —— 回来直接从下一节接着看。
🛑 读到这里可以停 —— 已经读了约 67 分钟。 最后一段还有(约 24 分钟):
gradcheck:不要靠眼睛检查导数 · 二阶导:默认能,除非你明说不能 · 这一章不讲什么 回来的时候不用重读,直接从下一节接着看就行。
🧪 五、gradcheck:不要靠眼睛检查导数
反向写错不会报错,只会让模型收敛得差一点 —— 这是本章所有坑里最贵的一个。
PyTorch 自带一个数值验证工具:它拿有限差分算一遍数值梯度,和你的 backward 对答案。
import torch
from torch.autograd import gradcheck
class Square(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return x * x
@staticmethod
def backward(ctx, g):
(x,) = ctx.saved_tensors
return g * 2 * x
class SquareWrong(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return x * x
@staticmethod
def backward(ctx, g):
(x,) = ctx.saved_tensors
return g * x # 💀 少乘了一个 2
x = torch.tensor([1.0, 2.0], dtype=torch.double, requires_grad=True) # ⭐ 必须 float64
print("正确的实现:", gradcheck(Square.apply, (x,)))
try:
gradcheck(SquareWrong.apply, (x,))
except Exception as e:
print("少乘 2 的:", type(e).__name__)
print(str(e)[:400])
实跑输出:
对照
正确的实现: True
少乘 2 的: GradcheckError
Jacobian mismatch for output 0 with respect to input 0,
numerical:tensor([[2.0000, 0.0000],
[0.0000, 4.0000]], dtype=torch.float64)
analytical:tensor([[1., 0.],
[0., 2.]], dtype=torch.float64)
⭐ numerical 是差分算的(可信),analytical 是你的 backward 算的。 这里对角线正好差 2 倍 —— 比值本身就是线索:差常数倍是系数写错,差符号是正负写反,只有某几个位置不对通常是 mask / 边界条件写错。
⚠️⚠️ gradcheck 必须用 float64
import torch
from torch.autograd import gradcheck
class Square(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return x * x
@staticmethod
def backward(ctx, g):
(x,) = ctx.saved_tensors
return g * 2 * x
x = torch.randn(4, requires_grad=True) # ⚠️ 默认 float32
gradcheck(Square.apply, (x,))
实跑输出(截断):
对照
UserWarning: Input #0 requires gradient and is not a double precision floating point or complex. This check will likely fail if all the inputs are not of double precision floating point or complex.
GradcheckError: Jacobian mismatch for output 0 with respect to input 0,
numerical:tensor([[-2.9802, 0.0000, 0.0000, 0.0000],
[ 0.0000, -1.2666, 0.0000, 0.0000],
💀 这个 Square 的实现是【对】的(上面刚验过 True),只是换成 float32 就挂了。
原因是有限差分要算 (f(x+h) - f(x-h)) / 2h,h 很小,两个相近的数相减在 float32 下几乎没有有效位剩下。
⚠️ 所以 gradcheck 失败的第一件事不是改代码,是先确认输入是不是 torch.double。 PyTorch 好心给了那条 warning,但它混在一堆输出里很容易被忽略。
🔁 六、二阶导:默认能,除非你明说不能
import torch
from torch.autograd.function import once_differentiable
class Square(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return x * x
@staticmethod
def backward(ctx, g):
(x,) = ctx.saved_tensors
return g * 2 * x
class SquareOnce(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return x * x
@staticmethod
@once_differentiable
def backward(ctx, g):
(x,) = ctx.saved_tensors
return g * 2 * x
x = torch.tensor([3.0], requires_grad=True)
g1 = torch.autograd.grad(Square.apply(x), x, create_graph=True)[0]
g2 = torch.autograd.grad(g1, x)[0]
print("一阶导 =", g1.item(), " 二阶导 =", g2.item())
x2 = torch.tensor([3.0], requires_grad=True)
g1b = torch.autograd.grad(SquareOnce.apply(x2), x2, create_graph=True)[0]
torch.autograd.grad(g1b, x2)
实跑输出:
要点
一阶导 = 6.0 二阶导 = 2.0
RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn
⭐ x² 在 x=3 处:一阶导 2x = 6,二阶导 2。两个都对。
这是因为 backward 里的 g * 2 * x 本身就是普通的可微算子,配上 create_graph=True(见第 02 章)它自然就能再求一次导。
⚠️ @once_differentiable 是你主动放弃这个能力:它在 backward 外面包了一层 no_grad,于是二阶导求不出来,报的还是那句 element 0 of tensors does not require grad。
✅ 什么时候该用它:你的 backward 里干了不可微的事(调 numpy、调外部库、写了 .item() 分支)。这时候主动挂上 @once_differentiable,让别人在第二次求导时收到一句明确的报错,好过让 autograd 在半路上给出一个悄悄错掉的结果。
💡 需要二阶导的场景比你想的多:MAML 这类元学习、WGAN-GP 的梯度惩罚、以及任何要算 Hessian-vector product 的地方。
🚦 七、这一章不讲什么
| 问题 | 去哪 |
|---|---|
一个 C++ / CUDA 算子怎么注册进 PyTorch(TORCH_LIBRARY、dispatcher、怎么编) |
《框架底下是 C++》05 · 算子怎么注册进 PyTorch ⭐ 本章只到 Python 这一层 |
| 激活占了显存的多少、重算该在哪里切、和 ZeRO/FSDP 怎么配 | 《AI基础设施》09 · 显存优化全家桶 |
| 注意力矩阵为什么可以不存、重算为什么反而更快 | 《AI基础设施》08 · FlashAttention |
| 反向传播本身的数学(链式法则怎么走) | 《机器学习与深度学习基础》08 · 反向传播 |
🔗 这一章连到哪里
| 相关的地方 | 为什么 |
|---|---|
| 02 · autograd:图是什么时候建的 | 本章第六节的 create_graph=True 出自那里;grad_fn 是什么、图为什么用完就扔,也都在那一章 |
| 03 · 就地操作:三种报错,三个不同的原因 | ⭐ 第三节那条 modified by an inplace operation … is at version 1 就是那一章第二种报错 —— 这里你看到了它是从哪一行代码里发出来的(save_for_backward 的版本检查) |
04 · detach / no_grad / requires_grad |
needs_input_grad[0] == False 的典型来源就是冻结了 backbone;重算那一节的 x.detach().requires_grad_(True) 也是那三个概念的组合用法 |
| 10 · 把模型交出去 | ⚠️ 自定义 Function 会直接影响下一章:导出的是推理图,你精心写的 backward 会被整个丢掉(对推理无所谓);但如果 forward 里绕出了框架(像本章的 numpy 例子),下一章实测导出会「成功」而结果是一个常数 |
| 《AI基础设施》08 · FlashAttention | ⭐ 那边讲「为什么重算能省显存」(算法与显存账),这一章讲「怎么把重算写进反向」(框架接口) —— 读完这里再去那里,recompute 那几段会突然变得具体 |
| 《AI基础设施》09 · 显存优化全家桶 | 梯度检查点在整套显存手段里排第几、和别的手段怎么叠 |
| 《AI基础设施》18 · 量化 | QAT 里那个「伪量化节点」的反向,就是第一节的 STE |
| 《框架底下是 C++》05 · 算子怎么注册进 PyTorch | ⭐ 本章的下一层:Python 侧的 Function 之外,一个真算子在 C++ 侧是怎么被 dispatcher 找到的 |
✅ 检查点
- 有人说「我这个
x.sigmoid() * w + b写成autograd.Function会更快」,哪里不对?真正需要自定义算子的是哪三种情况? torch.round的梯度实测是多少?为什么这个「正确答案」反而是个问题?STE 怎么解决它?Square()(x)报的是什么?为什么这条报错指错了方向?backward该返回几个值?forward收了一个张量和两个 float,backward怎么写?ctx.needs_input_grad是什么类型?第二节那个Clip例子里它实测是什么?ctx.save_for_backward(x)比ctx.x = x多做了什么?第三节那个实验里,两种写法分别得到什么结果?- 重算那一节的三个数字分别是多少?最后那个「最大差」为什么重要?
forward和backward里torch.is_grad_enabled()各是什么?由此推出哪两条写法上的结论?RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn在这一章出现了几次?分别是什么原因?gradcheck为什么必须用float64?它失败时第一件该做的事是什么?- 自定义
Function默认支不支持二阶导?@once_differentiable什么时候该主动加?
👀 答案
- 那个表达式全是已有算子的组合,autograd 自己求出来的导数又对又不慢,手写
Function只多一层 Python 调用和一个会写错的地方。真正需要的三种:①真实梯度是 0 或不存在(round、二值化、argmax)②前向绕出了框架(numpy / C++ / 外部库,图会断)③想用反向重算换显存。⚠️ 「我想让它更快」不在清单里。 - 实测
[0.0, 0.0, 0.0]——round是阶梯函数,除跳变点外导数处处为 0,框架没算错。问题是梯度为 0 意味着这一步之前的参数收不到任何信号,网络前半截等于没训。STE(直通估计)的做法是前向照样取整、反向假装是恒等函数:实测前向[0.0, 1.0, 1.0]、梯度[1.0, 1.0, 1.0],前向和反向故意不匹配。 RuntimeError: Legacy autograd function with non-static forward method is deprecated.(前面还有一条should not be instantiated的 DeprecationWarning)。⚠️ 它指错了方向:报错说forward不是 static,而那个forward明明写了@staticmethod;真正的原因是你实例化了(Square()而不是Square.apply)。看到non-static forward method先去找有没有多出来的括号。- 个数 =
forward输入参数的个数(ctx不算),一一对应。少了会报function ScaleBadBackward returned an incorrect number of gradients (expected 2, got 1)。一个张量 + 两个 float 就写return gx, None, None——None表示「这个入参没有梯度」,位置必须占住。 - 一个和输入参数等长的布尔元组。
Clip例子里实测是(True, False, False):只有x要梯度,lo/hi两个 float 不要。用它可以跳过没人要的那部分计算(比如冻结了 backbone 时)。 - 多做了版本检查:存的时候记下张量版本号,反向取的时候比对,对不上就报错。实测中前向存下的
b被b.mul_(10)就地改掉之后 ——ctx.x版没报任何错,算出[20.0, 40.0, 60.0](正确答案是[2.0, 4.0, 6.0],大了 10 倍,因为前向用旧值、反向读到新值);save_for_backward版抛RuntimeError: one of the variables needed for gradient computation has been modified by an inplace operation … is at version 1; expected version 0 instead. - 8.0 MB → 1.0 MB(512×512 float32 正好 1 MB,普通写法存 8 份中间结果、重算写法只存入口那 1 份)、
block()从 1 次变 2 次(省下的显存是用一次额外前向换的)、最大差0.00e+00。最后这个重要是因为它说明重算不是近似方法,梯度逐位相同,算的是同一件事。 - 实测
forward里是 False、backward里也是 False(外面是 True)。结论:①forward里不需要自己写with torch.no_grad(),它本来就不建图;② 重算那一段必须写with torch.enable_grad():,否则重跑出来的结果没有grad_fn。 - 三次,三个不同的原因:① 第一节 numpy 那段 —— 中途绕出框架,
grad_fn是None,图断了;② 第四节忘了enable_grad——backward里 grad 默认是关的;③ 第六节@once_differentiable之后求二阶导 —— 是主动放弃的结果。这句话的含义永远是「你要求导的东西不在图上」,但成因不止一个(还有最常见的「忘了requires_grad=True」)。 - 因为有限差分要算
(f(x+h) - f(x-h)) / 2h,两个相近的数相减在float32下几乎没有有效位剩下。实测一个正确的Square实现在float64下返回True、换成float32就抛GradcheckError,同时给出UserWarning: Input #0 requires gradient and is not a double precision floating point or complex.✅ 失败时第一件事是确认输入是不是torch.double,不是改代码。失败信息里numerical可信、analytical是你写的,两者的比值是线索(实测差 2 倍就是少乘了个 2)。 - 默认支持 —— 因为
backward里的g * 2 * x本身就是可微算子,配create_graph=True就能再求一次(实测x=3时一阶导 6.0、二阶导 2.0)。@once_differentiable是主动放弃它(内部包了一层no_grad)。✅ 该加的时候:backward里干了不可微的事(调 numpy、调外部库、.item()分支)—— 让别人求二阶导时收到明确报错,好过拿到一个悄悄错掉的结果。
🛑 可以停在这里
⚡ 走神救援
⭐
autograd.Function是一份合同:前向随便你怎么算,代价是反向那一格必须你自己填。真正需要它的只有三种:真实梯度是 0 或不存在(取整的梯度确实是 0,⭐ 框架没算错,但这意味着前半个网络收不到任何信号——于是有了直通估计器:前向取整、反向假装恒等,故意让两边不匹配,量化感知训练就靠它)、前向绕出了框架(用别的库算完再回来,那条链就断了,⭐ Function 就是把断口焊回去的东西)、拿重算换显存。⚠️ 「我想让它更快」不在这份清单里。
⚠️ 最容易踩的一条:必须走
.apply()调用——写成普通调用时的报错指错了方向,真正的原因是你把它实例化了。💀 最贵的一节是「把张量直接挂在 ctx 上」:它只存引用,前向存下的张量被就地改掉之后,⭐ 它不报错,只是算出一个错得离谱的梯度(前向用旧值、反向用新值)。而专用的保存接口多记了一个版本号,同样场景会直接抛出「被就地操作修改过」。⭐ 规则:张量一律用专用接口保存,非张量才挂 ctx。
♻️ 重算换显存的实测很干净:只存入口、反向重跑一遍,显存降一个量级,代价是那段算两次,⭐ 而梯度逐位相同、不是近似。这个 Function 框架已经写好了,实际项目直接用现成的。
🧪 ⭐ 梯度检查必须用双精度:同一个正确的实现换成单精度照样失败(有限差分在单精度下没有有效位)——⚠️ 失败第一件事是查 dtype,不是改代码;真错了的时候,两个梯度的比值就是线索。
⚠️ 最后一个坑:前向里梯度是关着的(不用自己写),⭐ 而反向里也是关着的——要二阶导就必须显式打开。
下一节 👉 10-把模型交出去.md