如何获取PyTorch中backward()的grad_fn可用函数完整列表及相关公式?
PyTorch中获取所有xxxBackward0类及对应正反向公式的方法
一、列出所有xxxBackward0类
PyTorch没有直接的命令一次性列出所有这类反向传播类,但可以通过以下两种方式筛选收集:
1. 遍历torch.autograd.function模块
反向传播类大多继承自torch.autograd.Function,可通过反射遍历模块属性筛选:
import torch import inspect backward_classes = [] # 遍历模块内所有成员,筛选以Backward0结尾的类 for name, obj in inspect.getmembers(torch.autograd.function): if inspect.isclass(obj) and name.endswith('Backward0'): backward_classes.append((name, obj)) # 打印结果 for name, cls in backward_classes: print(f"反向类: {name}")
注:这种方式可能遗漏部分分散在其他模块的反向类。
2. 通过正向操作触发反向类生成
更可靠的方式是执行常见正向操作,获取对应grad_fn并收集:
import torch # 定义常见张量操作集合 operations = [ (torch.add, "加法"), (torch.mul, "乘法"), (torch.sub, "减法"), (torch.matmul, "矩阵乘法"), (torch.pow, "幂运算"), (torch.sin, "正弦"), (torch.cos, "余弦"), (torch.div, "除法"), ] backward_info = {} for op, op_name in operations: # 创建可求导张量 a = torch.tensor([1.0], requires_grad=True) b = torch.tensor([2.0], requires_grad=True) if op not in [torch.sin, torch.cos] else None # 执行正向操作 out = op(a, b) if b is not None else op(a) # 记录反向类及对应正向信息 grad_fn = out.grad_fn if grad_fn is not None: cls_name = grad_fn.__class__.__name__ backward_info[cls_name] = { "正向操作": op_name, "正向表达式": f"out = {op_name}(a, b)" if b is not None else f"out = {op_name}(a)" } # 打印收集到的信息 for cls_name, info in backward_info.items(): print(f"\n反向类: {cls_name}") print(f"对应正向操作: {info['正向操作']}") print(f"正向表达式: {info['正向表达式']}")
二、获取正反向计算公式
PyTorch没有内置命令直接输出反向公式,可通过以下方式获取:
- 查看反向类源码:用
inspect.getsource(grad_fn.__class__)查看反向类实现代码,推导梯度计算逻辑。比如AddBackward0的反向逻辑是将上游梯度直接传递给两个输入张量。 - 手动推导:基于链式法则推导基础操作的反向公式,比如乘法操作的反向梯度是上游梯度分别乘以两个输入张量。
常见示例:
- AddBackward0:
- 正向公式:
out = a + b - 反向公式:
∂out/∂a = 1,∂out/∂b = 1→ 上游梯度直接传递给a和b
- 正向公式:
- MulBackward0:
- 正向公式:
out = a * b - 反向公式:
∂out/∂a = b,∂out/∂b = a→ 上游梯度分别乘以b和a后传递
- 正向公式:
- SubBackward0:
- 正向公式:
out = a - b - 反向公式:
∂out/∂a = 1,∂out/∂b = -1→ 上游梯度直接给a,取反后给b
- 正向公式:
内容的提问来源于stack exchange,提问作者user13474103
相关产品推荐
相关产品推荐

