首页 / PyTorch 入门教程 / FX:图变换与图优化

PyTorch 入门教程

FX:图变换与图优化

本教程共 60 篇 · 第 48 篇 · 更新于 2026-08-17 · 约 4 分钟阅读

PyTorchFX图变换symbolic_traceInterpreter图优化

本节目标:认识 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.linearcall_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 能在调用子模块前后插手。想统计某类操作,覆写对应钩子就行。官方教程里的分析器覆写了 runrun_node,分别记录整网和单节点耗时。用 time.time() 计时精度一般,真实性能分析建议用第 32 章的 Profiler。Interpreter 的价值在于灵活:你想统计什么都能自己写。

Note

Interpreter 适合分析和教学场景。生产环境跑模型还是用 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 全部替换成 mulx + y * x 变成了 x * y * x。因为替换发生在编译期,运行时没有任何额外开销。适合做算子替换、调试,或者给特定后端接自定义实现。

它和 FX 的分工是:FX 改的是图的结构,Mode 改的是算子分发的行为。两者还能配合,Mode 在追踪期改写,FX 在图上继续变换。一个典型场景:某个算子在新版本里数值行为变了,你想全项目临时换回旧行为,Mode 一段代码就搞定,不用到处改调用点。

Note

Torch Function Mode 需要 PyTorch 2.7 以上,本章按 2.13 基线没问题。普通(未编译)代码里也能用,但每次算子调用都有模式分发开销,和 torch.compile 搭配才划算。命名里的 Base 前缀表示它接管所有 torch.* 算子,写一个类就全覆盖。

小结

FX 提供了「看得到、改得动」的计算图:symbolic_trace 抓图,节点类型帮你理解结构,遍历改写实现图变换,Interpreter 做自定义分析。下一章开始进入导出与部署,torch.export 和 ONNX 都用得上图的思想。这一章学的节点、改写、解释执行,全是后面内容的地基,值得多花几分钟好好消化。