FX:图变换与图优化
本教程共 60 篇 · 第 48 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:认识 PyTorch 的 FX 子系统,学会把模型变成可检查、可改写的计算图,并用 Interpreter 做自定义分析。
什么是 FX
FX 是 PyTorch 自带的「图」工具包,全名 torch.fx。它能把一段 Python 代码抓成计算图(graph):每个操作是一个节点(node),节点之间用数据依赖连成一张网。有了图,就能程序化地检查、分析和改写模型。
打个比方。你的模型 forward 像一份做菜步骤,eager 模式是照着步骤一步步做。FX 先把步骤抄成一张流程图贴在墙上,你可以端详、可以改,改完再照着做。
FX 在 PyTorch 内部无处不在:量化工具、torch.compile 的抓图环节,底层都是 FX。学它不亏。
FX 从 PyTorch 1.8 起随框架发布,接口一直标着 beta,但用量非常大。它解决的核心问题:让模型结构变成可编程的数据。没有图的时候,你想批量改模型只能手改每个模块;有了图,一段循环就能改遍全网。举个具体的:想给所有卷积加个调试打印,手写要改十几处,FX 一条循环遍历节点就搞定。这就是图的价值,也是后面所有图工具的共同起点。
抓图:symbolic_trace
核心 API 是 torch.fx.symbolic_trace。它用假张量跑一遍 forward,把真实发生的操作记成图:
import torch
import torch.fx
import torch.nn.functional as F
class Net(torch.nn.Module):
def __init__(self):
super().__init__()
self.linear = torch.nn.Linear(8, 8)
def forward(self, x):
return F.relu(self.linear(x))
net = Net()
gm = torch.fx.symbolic_trace(net)
print(gm.graph)
输出大概是:
graph():
%x : [num_users=1] = placeholder[target=x]
%linear : [num_users=1] = call_module[target=linear](args = (%x,), kwargs = {})
%relu : [num_users=1] = call_function[target=torch.nn.functional.relu](args = (%linear,), kwargs = {})
return relu
gm 是 GraphModule,一个「图 + 模块」的合体。它还能像普通模型一样调用:gm(x) 和 net(x) 结果一致。内部实现上,symbolic_trace 用一个特殊代理对象跑 forward,任何张量操作都被记录下来,而不是真的计算。所以 trace 很快,几毫秒就完成。
图里节点有几种固定类型:
placeholder:函数入参。call_module:调用子模块,比如linear。call_function:调用普通函数,比如relu。get_attr:读取参数。output:返回值。
placeholder 对应函数参数;call_module 的 target 是子模块名,比如 linear 指向 self.linear;call_function 的 target 是函数对象本身;get_attr 读取 self 上的参数,比如权重。遍历 gm.graph.nodes 就能拿到全图信息,每个节点用 args 引用上游节点,形成数据流。图里的每条边,本质就是上游到下游的数据传递。
看懂节点类型,就能写变换了。
改写图:把 ReLU 换成别的
图的价值在于可改写。遍历所有节点,改 target 就能换算子:
for node in gm.graph.nodes:
if node.target == torch.nn.functional.relu:
node.target = torch.sin # 把 ReLU 换成 sin
gm.recompile()
x = torch.randn(4, 8)
print(gm(x)) # 现在算的是 sin(linear(x))
recompile() 根据改好的图重新生成 Python 代码。torch.fx 内置了一批现成变换,官方教程里有个典型的例子:用 FX 写 Conv-BN 融合器——找到卷积和批归一化两个节点,把它们合并成一个算子,推理时少算一次。
FX 的量化路径 torch.ao.quantization 也是同样的思路:先抓图,再插入量化、反量化节点。
替换之外,还能插入和删除节点:graph.inserting_before(node) 在指定位置前插新节点,node.replace_all_uses_with(新节点) 把下游引用全部改指。前面说的 Conv-BN 融合,就是靠这些 API 组合出来的。这些 API 组合起来,能写的变换几乎不受限。变换写完记得 recompile(),不然改的是图,跑的还是旧代码。
Warning
symbolic_trace用假张量跑forward,依赖真实数据值的代码它处理不了。比如if x.sum() > 0这种分支,trace 时走哪条由假张量决定,结果可能不对。控制流复杂的模型,得用torch.cond显式表达,或者交给第 50 章的torch.export。FX 适合结构规整的网络。
Interpreter:逐节点执行图
GraphModule 是把图编译成代码跑。另一种跑法叫 Interpreter:不编译,逐个节点解释执行。好处是每执行一个节点,你都能插手。
继承 Interpreter,覆写 run_node,就能做自定义分析。下面写一个极简性能分析器,统计每个节点耗时:
import time
import torch
import torch.fx
class TimingInterpreter(torch.fx.Interpreter):
def __init__(self, module):
gm = torch.fx.symbolic_trace(module)
super().__init__(gm)
self.times = {}
def run_node(self, node):
start = time.time()
result = super().run_node(node)
self.times.setdefault(str(node), []).append(time.time() - start)
return result
interp = TimingInterpreter(net)
interp.run(torch.randn(4, 8))
for name, times in interp.times.items():
print(f"{name}: {sum(times) / len(times) * 1000:.2f} ms")
官方教程用这个思路分析了 ResNet18,结论是 MaxPool2d 最耗时。Interpreter 的执行是分层的:run 负责整体流程,run_node 负责单个节点,节点内部还有更细的钩子,比如 run_call_module 能在调用子模块前后插手。想统计某类操作,覆写对应钩子就行。官方教程里的分析器覆写了 run 和 run_node,分别记录整网和单节点耗时。用 time.time() 计时精度一般,真实性能分析建议用第 32 章的 Profiler。Interpreter 的价值在于灵活:你想统计什么都能自己写。
NoteInterpreter 适合分析和教学场景。生产环境跑模型还是用 GraphModule 或
torch.compile,直接编译成内核,快得多。
除了计时,Interpreter 还能做更多:检查中间结果、给特定层喂替换数据、模拟剪枝效果。凡是需要「跑一步看一步」的场景,它都合适。
Torch Function Mode:追踪时改写算子
FX 改的是「图上的节点」。还有个更轻的玩法:Torch Function Mode。它拦截所有 torch.* 算子调用,在追踪(trace)时直接替换实现。配合 torch.compile 使用,改写行为零运行时开销。
import torch
from torch.overrides import BaseTorchFunctionMode
class AddToMultiplyMode(BaseTorchFunctionMode):
def __torch_function__(self, func, types, args=(), kwargs=None):
if func == torch.Tensor.add:
func = torch.mul
return super().__torch_function__(func, types, args, kwargs)
@torch.compile
def test_fn(x, y):
return x + y * x
x = torch.rand(2, 2)
y = torch.rand_like(x)
with AddToMultiplyMode():
z = test_fn(x, y)
assert torch.allclose(z, x * y * x) # 加法真的变成了乘法
这段代码把 add 全部替换成 mul,x + y * x 变成了 x * y * x。因为替换发生在编译期,运行时没有任何额外开销。适合做算子替换、调试,或者给特定后端接自定义实现。
它和 FX 的分工是:FX 改的是图的结构,Mode 改的是算子分发的行为。两者还能配合,Mode 在追踪期改写,FX 在图上继续变换。一个典型场景:某个算子在新版本里数值行为变了,你想全项目临时换回旧行为,Mode 一段代码就搞定,不用到处改调用点。
NoteTorch Function Mode 需要 PyTorch 2.7 以上,本章按 2.13 基线没问题。普通(未编译)代码里也能用,但每次算子调用都有模式分发开销,和
torch.compile搭配才划算。命名里的 Base 前缀表示它接管所有torch.*算子,写一个类就全覆盖。
小结
FX 提供了「看得到、改得动」的计算图:symbolic_trace 抓图,节点类型帮你理解结构,遍历改写实现图变换,Interpreter 做自定义分析。下一章开始进入导出与部署,torch.export 和 ONNX 都用得上图的思想。这一章学的节点、改写、解释执行,全是后面内容的地基,值得多花几分钟好好消化。