必须自定义 torch.autograd.Function 是因为 torch.nn.functional 缺少特定梯度逻辑(如 STE 中的梯度截断或替换),需手动实现 forward 和 backward;参数顺序严格对应,中间变量须用 ctx.save_for_backward() 保存,非 Tensor 输入需存为 ctx 属性,且 backward 返回梯度元组须与 forward 输入数量一致。

为什么不能直接用 torch.nn.Functional 而要写自定义 torch.autograd.Function
因为 torch.nn.functional 里没有你需要的梯度逻辑——比如你想让反向传播中某个梯度被截断、缩放、或替换成固定值(如 STE 中的 sign 函数),又或者你封装了一个外部 C++/CUDA 算子但没提供自动求导支持。这时候必须继承 torch.autograd.Function,手动实现 forward 和 backward。
forward 和 backward 的参数怎么对齐?
关键点是:forward 的输入参数(除 cls 外)会原样传给 backward,但 backward 的第一个参数是输出梯度 grad_output,其余才是 forward 的输入——且顺序严格一致。中间变量要用 ctx.save_for_backward() 显式保存,否则反向时拿不到。
-
forward返回值必须是 Tensor 或 Tensor 元组,不能是 Python 数字或 NumPy 数组 -
backward必须返回与forward输入数量相同的梯度元组,对应每个输入;不需要梯度的位置填None - 如果
forward输入含非 Tensor(如int、bool),它们不会参与求导,也**不能**传给backward——得在ctx上存为属性,例如ctx.threshold = threshold
如何避免“梯度不流动”或“RuntimeError: Trying to backward through the graph a second time”?
常见原因是:在 forward 中误用了 .detach()、.data 或 torch.no_grad();或者把需要梯度的 Tensor 当作普通 Python 对象处理(比如放进 list 再 pop 出来)。另外,backward 返回的梯度 shape 必须和对应输入完全一致,否则会静默失败或报错。
- 所有参与计算的 Tensor 都要确保
requires_grad=True(除非你明确想切断梯度) - 不要在
forward中调用.item()、.numpy()、.cpu()等脱离计算图的操作 - 调试时可在
backward开头加print(f"grad_output shape: {grad_output.shape}"),确认输入输出维度匹配 - 若函数需支持高阶导数(如用于 Hessian 计算),要在
@staticmethod前加@once_differentiable装饰器(默认不支持)
一个带 mask 的 STE 示例(用于二值网络)
这是最常被查的场景:前向用 torch.sign(x),反向却希望梯度走 torch.clamp(1 - x**2, min=0)(即直通估计器)。
立即学习“Python免费学习笔记(深入)”;
class SignSTE(torch.autograd.Function):
@staticmethod
def forward(ctx, x):
ctx.save_for_backward(x)
return torch.sign(x)
<pre class='brush:python;toolbar:false;'>@staticmethod
def backward(ctx, grad_output):
x, = ctx.saved_tensors
grad_input = grad_output * torch.clamp(1 - x ** 2, min=0)
return grad_input
注意:这里没做 ctx.mark_non_differentiable(),因为输出本身是可微的(只是我们覆盖了它的梯度);如果输出含不可微成分(如索引、布尔判断),才需要标记。
真正容易被忽略的是:这个函数不能直接当模块用,必须显式调用 SignSTE.apply(x);如果想封装成 nn.Module,得在 forward 方法里调用它,而不是继承 Function 自己。


















