🏠 总目录📚 本教程 09 · 自定义算子 ← →
📑 本页目录(点开跳转)

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

⭐ 两条都反直觉,各解决一个问题:

  1. forward 里 grad 已经是关的。 所以你在 forward 里做什么都不会建图 —— 不需要自己加 with torch.no_grad()(加了也无害,只是多余)。这也解释了第二节那个 numpy 例子为什么不用担心图被污染。
  2. ⚠️ 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 找到的

✅ 检查点

  1. 有人说「我这个 x.sigmoid() * w + b 写成 autograd.Function 会更快」,哪里不对?真正需要自定义算子的是哪三种情况?
  2. torch.round 的梯度实测是多少?为什么这个「正确答案」反而是个问题?STE 怎么解决它?
  3. Square()(x) 报的是什么?为什么这条报错指错了方向?
  4. backward 该返回几个值?forward 收了一个张量和两个 float,backward 怎么写?
  5. ctx.needs_input_grad 是什么类型?第二节那个 Clip 例子里它实测是什么?
  6. ctx.save_for_backward(x) 比 ctx.x = x 多做了什么?第三节那个实验里,两种写法分别得到什么结果?
  7. 重算那一节的三个数字分别是多少?最后那个「最大差」为什么重要?
  8. forward 和 backward 里 torch.is_grad_enabled() 各是什么?由此推出哪两条写法上的结论?
  9. RuntimeError: element 0 of tensors does not require grad and does not have a grad_fn 在这一章出现了几次?分别是什么原因?
  10. gradcheck 为什么必须用 float64?它失败时第一件该做的事是什么?
  11. 自定义 Function 默认支不支持二阶导?@once_differentiable 什么时候该主动加?
👀 答案
  1. 那个表达式全是已有算子的组合,autograd 自己求出来的导数又对又不慢,手写 Function 只多一层 Python 调用和一个会写错的地方。真正需要的三种:①真实梯度是 0 或不存在(round、二值化、argmax)②前向绕出了框架(numpy / C++ / 外部库,图会断)③想用反向重算换显存。⚠️ 「我想让它更快」不在清单里。
  2. 实测 [0.0, 0.0, 0.0] —— round 是阶梯函数,除跳变点外导数处处为 0,框架没算错。问题是梯度为 0 意味着这一步之前的参数收不到任何信号,网络前半截等于没训。STE(直通估计)的做法是前向照样取整、反向假装是恒等函数:实测前向 [0.0, 1.0, 1.0]、梯度 [1.0, 1.0, 1.0],前向和反向故意不匹配。
  3. 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 先去找有没有多出来的括号。
  4. 个数 = forward 输入参数的个数(ctx 不算),一一对应。少了会报 function ScaleBadBackward returned an incorrect number of gradients (expected 2, got 1)。一个张量 + 两个 float 就写 return gx, None, None —— None 表示「这个入参没有梯度」,位置必须占住。
  5. 一个和输入参数等长的布尔元组。Clip 例子里实测是 (True, False, False):只有 x 要梯度,lo / hi 两个 float 不要。用它可以跳过没人要的那部分计算(比如冻结了 backbone 时)。
  6. 多做了版本检查:存的时候记下张量版本号,反向取的时候比对,对不上就报错。实测中前向存下的 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.
  7. 8.0 MB → 1.0 MB(512×512 float32 正好 1 MB,普通写法存 8 份中间结果、重算写法只存入口那 1 份)、block() 从 1 次变 2 次(省下的显存是用一次额外前向换的)、最大差 0.00e+00。最后这个重要是因为它说明重算不是近似方法,梯度逐位相同,算的是同一件事。
  8. 实测 forward 里是 False、backward 里也是 False(外面是 True)。结论:① forward 里不需要自己写 with torch.no_grad(),它本来就不建图;② 重算那一段必须写 with torch.enable_grad():,否则重跑出来的结果没有 grad_fn。
  9. 三次,三个不同的原因:① 第一节 numpy 那段 —— 中途绕出框架,grad_fn 是 None,图断了;② 第四节忘了 enable_grad —— backward 里 grad 默认是关的;③ 第六节 @once_differentiable 之后求二阶导 —— 是主动放弃的结果。这句话的含义永远是「你要求导的东西不在图上」,但成因不止一个(还有最常见的「忘了 requires_grad=True」)。
  10. 因为有限差分要算 (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)。
  11. 默认支持 —— 因为 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

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