解密PostGradPassManager:自定义融合为何总失效?
2026/9/19 3:05:03 网站建设 项目流程

跑过几轮torch.compile之后,很多人都会遇到一个有点魔幻的场面:手写了一个自定义 fusion,信心满满地塞进编译流程,结果一看生成的 kernel,跟优化前一模一样,连个影子都没有。别急着怀疑编译器后端,先静下来问自己一个问题:这个 pass 到底挂在了图变换流水线的哪个阶段?绝大多数 fusion 不生效,不是代码写错,而是挂错了阶段,而PostGradPassManager就是最容易踩中的那个位置。

这一篇我会先把图变换流水线的整体结构说清楚,讲明白图从哪来、要到哪去、每个阶段在解决什么问题,然后深入PostGradPassManager的职责边界,带你看懂它为什么叫"post grad",以及它跟前 grad、后端的边界到底划在哪。之后我会手写一个真正能跑的自定义 fusion:从 FX 子图匹配到 Inductor lowering,再到 IR 转储验证,把整条链路跑通。内容会以 PyTorch 2.x 的 torch.compile 路径为背景,同时顺带对比传统编译器里PassManager的设计理念,适合已经写过一点算子融合、但对整个编译流程还停留在黑盒阶段的同学。

1. 图变换流水线到底在编什么:从计算图到“能跑的代码”

1.1 计算图不是一种,而是三种

很多人以为图变换流水线只处理一张图,这是最大的误解。一次完整的编译,从模型代码到最终能在 GPU 上跑的 kernel,中间至少会经历三种形态完全不同的图。这三种图之间不是简单的等价转换,而是在不同的抽象层级上反复重写。

第一层是TorchDynamo 从字节码层面捕获的 FX Graph。这张图跟 Python 执行关系非常紧密,里面可能还残留着getattrassert、Python 控制流留下的子图调用。它的职责是把你写的 Python 代码"翻译"成算子调用序列,但远没有到适合做优化的程度。比如你写一个循环,Dynamo 可能捕获成多个call_function节点,也可能整体变成一个 subgraph,里面继续走 Python 解释器,这在后续优化里都非常烫手。

第二层是AOTAutograd 产出的、已经展开自动求导的 Graph。这一层的图里不再有autograd.Function那样隐含的求导逻辑,而是显式地把 forward 和 backward 需要的算子都展开成一张大图。也正因为求导被展开了,很多原来被隐藏在 autograd 黑盒里的冗余计算、可以被消掉的中间变量,现在都暴露出来了。这个阶段产出的图,才真正适合做代数化简、公共子表达式消除、死代码消除等操作。

第三层是Inductor IR。FX Graph 仍然是算子级别的抽象,而 Inductor 会把算子再拆成更细粒度的循环、buffer、布局描述。它不再是"图"的形态,而是一组PointwiseReductionTemplateBuffer等 IR 节点。到了这一层,PassManager 的职责开始转向循环级变换、调度顺序、缓存局部性这类底层问题。

所以当大家说"图变换流水线"的时候,一定要先确认自己站在哪一层。每一层的图结构不同,数据依赖的表达方式不同,能做和不能做的变换也完全不同。

1.2 为什么必须多出一个 PostGrad 阶段

不同编译器对"图变换流水线"的划分很不一样。LLVM 里是ModulePassManagerFunctionPassManager,MLIR 里有从ModuleFunc再到Op的多级 Pipeline。而 PyTorch 里的划分逻辑,核心看的是"自动求导的边界"。

TorchDynamo 捕获到的图,在 AOTAutograd 展开之前,其实还带着 forward/backward 的隐式结构。如果在此时直接对整张图做融合,会碰上一个很尴尬的问题:你并不知道哪些节点属于 forward 里需要为 backward 保留中间结果的节点,哪些是纯粹可丢弃的临时量。一旦你贸然把某个中间结果融合进后面一个大 kernel,而 backward 仍然单独引用它,这一步融合就打破了反向传播的依赖关系,轻则导致重复计算,重则直接产出错误的梯度。

所以从设计上就必须把"自动求导展开"作为一道分水岭。展开之后,forward 和 backward 之间的依赖全部变成了显式的数据依赖,图上的每个节点都知道谁在用它的输出,哪些中间结果可以重算,哪些必须保留。这个展开后的阶段,就是PostGradPassManager所处的窗口。

这也是它名字的由来:这里的 "grad" 不是指梯度,而是指位于grad-enabled autograd graph 之后的阶段。它处理的是已经完成自动求导展开、不再携带任何 autograd 上下文、只由底层算子组成的纯净计算图。这一步跨过去之后,才能开始放心大胆地做融合和重排。

2. PostGradPassManager 的职责边界:它管哪一段,不碰哪一段

2.1 自动求导展开后的“后处理”到底在做什么

PostGradPassManager在实际编译路径里的位置,大致是compile_fx函数内部召集的一连串针对后向展开图的 pass 集合。它接收的是一张 FX Graph,输出的仍然是一张 FX Graph,但是内容上已经从"刚展开的原始图"变成了"经过多轮化简和融合的优化图"。

这些 pass 做的事情大致可以分成四类。

一类是移除无意义节点。自动求导展开后经常产生很多viewclonedetachalias这类节点,它们的存在往往只是为了让 autograd 的链式求导规则能够成立。在图变换阶段它们已经没有任何价值,留下来只会干扰后续 pattern 匹配。因此会有一批专门做 noop 消除的 pass 把它们删掉。

第二类是代数化简与维度整理。比如把split之后再cat这种逆操作消除,把多次unsqueeze/squeeze合并成一次,把permute链化简成单个permute。这些操作在展开后的图上特别常见,因为反向传播里为了对齐梯度 shape,会大量插入 reshape 类算子。

第三类是pattern 识别式融合。像fuse_attentionfuse_conv_bnfuse_transpose_matrix_multiplication这类 pass,本质上是拿一个事先定义好的子图模板去匹配当前图,一旦命中,就把这一整片替换成一个更高层的算子或一个自定义 IR 节点。这也是大多数"自定义 fusion"应当挂靠的层面。

第四类是调度和后端相关的重写。例如把某些算子从 compute-intensive 改成 memory-bound 的执行策略,或者根据后端特性和 buffer 布局调整算子输入输出的排布。这类变换通常已经接近 IR 层,但一部分决策仍然发生在 FX 层。

2.2 Pass 不是想插哪就插哪:注册阶段决定生死

这是整个图变换流水线里最容易让人翻车的地方。很多人写了一个自定义 fusion pass,直接挂在torch._inductor.config.post_grad_custom_pre_pass或者手动插进post_grad_passes列表,结果发现某些子图怎么都匹配不上。原因往往不是你的匹配逻辑写错了,而是你选错了运行阶段。

一个自定义 fusion 可以有多个合理的插入点,每个点的语义和后续行为完全不同:

插入阶段输入图形态适合做什么容易踩的坑
Dynamo 捕获后、AOT 前带 Python 语义的 FX Graph粗粒度图改写、利用 Python 条件信息自动求导展开后你的 pattern 可能整个散掉
AOTAutograd 内部展开前/后的中间状态处理与自动求导相关的特殊语义API 变化快,几乎每个版本都在改
PostGradPassManager(推荐)自动求导展开后的纯净图算子层面 fusion、代数化简、pattern 替换pass 之间隐含顺序依赖,插错位置等于没跑
Inductor scheduler 层IR 节点循环体循环融合、布局优化已经是底层 IR,改起来成本高,调试困难

PostGradPassManager之所以是自定义 fusion 的主战场,是因为它刚好卡在"求导结构已消失"和"后端代码生成未开始"之间。到了这个阶段,你能看到的所有算子都已经是实际要计算的算子,不再有 Python 控制流和 autograd 包装,pattern 匹配的确定性最高。但与此同时,你也必须尊重这个阶段已经形成的 pass 顺序。

给一个非常直接的结论:如果你的 fusion pass 需要依赖某些已经做过的化简(比如 view 消除、transpose 合并),那就把 pass 放在这些化简 pass 之后;如果你的 fusion pass 会产生新节点,而后面的 pass 不一定认识这些新节点,那最好在 pass 内自己把后续影响一并处理好。我的习惯是优先使用post_grad_custom_pre_pass这类官方预留的钩子,而不是直接改post_grad_passes列表源码,因为前者在版本升级时兼容性更好,后者几乎每次 PyTorch 小版本更新都要跟着修。

3. 手写一个自定义 fusion:以 softmax 子图匹配为例

3.1 先定融合边界:哪些节点能进同一个 kernel

理论讲再多,不如亲手把一个 fusion 写通。我选一个既简单又典型的例子:把"exp(x) / sum(exp(x), dim=-1, keepdim=True)"这种 custom softmax-like 子图,融合成一个自定义算子。为什么要选这个?因为它包含了两种最常见的节点类别:pointwise 的exp和 reduction 的sum。在真正写 fusion 之前,你先要自己回答一个后端问题:这种子图到底适不适合融合成一个 kernel?

如果按 naive 方式拆开执行,exp要写出一个中间张量写回全局内存,sum要再去读这个中间张量做归约,然后再用div做第二个 pointwise。一次往返意味着多次全局内存读写。融合的核心收益正是把中间张量按 tile 留在寄存器或者共享内存里,让 reduction 和 pointwise 在同一个循环体里完成。这就是值得做的融合。

但这里还要注意一个问题:exp(x)的结果被两处使用,一处是sum,另一处是div。如果融合成一个算子,意味着exp(x)的结果要么被保留成中间 buffer,要么在同一个 kernel 里被计算两次。真正的 Triton 融合通常会选择保留中间值到一个中间 buffer,再在后续循环里复用它,因为重复计算指数函数的代价比读一次中间 buffer 更高。这个小决策就属于"融合边界"的范畴,决定你这个 kernel 到底是省了带宽还是反而增加了计算量。

3.2 在 FX 图上做 pattern 匹配并替换

确定边界之后,第一步是在 FX 图上把这个 pattern 找出来。我用一个简化版本演示:匹配exp(x) -> sum(y, dim=-1, keepdim=True) -> div(y, s)这条路径,并把它替换成自定义算子torch.ops.example.softmax_like(x, dim=-1)

下面的代码是目前 PyTorch 2.x 里常用的遍历式匹配方式。虽然标准库也有SubgraphMatcher,但它属于内部 API,不同版本位置和签名变化很大,所以我更推荐在代码里显式控制匹配逻辑,这样可控性更强,也更容易调试:

import torch import torch.fx as fx def match_softmax_like(graph: fx.Graph): matches = [] for div_node in graph.nodes: if div_node.op != "call_function" or div_node.target != torch.div: continue sum_node, y_node = div_node.args if sum_node.op != "call_function" or sum_node.target != torch.sum: continue if y_node.op != "call_function" or y_node.target != torch.exp: continue if y_node not in sum_node.args and y_node not in sum_node.kwargs.values(): continue # 进一步检查 dim 和 keepdim 参数,这里省略部分防护逻辑 matches.append((div_node, sum_node, y_node)) return matches

这里的代码故意写得很直白,目的是让你看清楚 pattern 匹配的本质:它就是在 DAG 上做子图同构查找。真正实战时还需要补充很多细节,比如必须验证sumdim参数是尾部维度、keepdim=Truediv_node的第二个参数必须是那个 sum 输出而不是别的标量,还要检查这些节点的使用计数,确保y_node除了被sumdiv使用之外没有被其他地方引用。

匹配到之后,替换逻辑并不复杂,但有一个容易忽略的坑:不能直接graph.erase_node旧节点,得先创建新节点,把旧节点的所有使用点替换到新节点上,然后再按依赖顺序从后往前删除。FX 没有自动管理这个引用关系,顺序没处理好会直接报"node has no users"或者更诡异的错误。

def replace_with_fused_op(graph: fx.Graph, matches): for div_node, sum_node, y_node in matches: x_node = y_node.args[0] new_node = graph.call_function( torch.ops.example.softmax_like, (x_node, -1), ) div_node.replace_all_uses_with(new_node) graph.erase_node(div_node) graph.erase_node(sum_node) graph.erase_node(y_node)

3.3 把自定义节点送进 Inductor 的 lower 流程

替换成torch.ops.example.softmax_like只是第一步,真正决定它能不能生成高效代码的是 lowering 阶段。Inductor 在收到一张带自定义算子的 FX Graph 时,会到lowerings这个注册表里找这个算子的实现。如果你没有注册任何 lowering,它通常会 fallback 到一个 eager 调用,也就是说你的 fusion"合法了",但一点也没有优化到。

自定义 lowering 有两种路线,取决于你想要多深的控制。

轻量路线是注册一个@torch.library.impl或直接注册到 Inductor 的 lowering 表,让自定义算子被拆成一个或者若干个 Inductor 原生 IR 节点。比如你可以把softmax_like拆成一个PointwiseIR 节点和一个ReductionIR 节点,Inductor 的 scheduler 会自动在循环层面继续尝试融合它们。这种做法的好处是能利用 Inductor 已有的调度优化,坏处是你无法精确控制最终生成的循环布局。

重量路线是自己定义一个TritonTemplate,完全控制 kernel 的循环逻辑、tile 大小和向量化方式。这个代码量会大很多,但才是真正意义上"写了一个自定义 fusion kernel"。最典型的写法是这样的骨架:

from torch._inductor import config from torch._inductor.codegen.triton_utils import signature_of from torch._inductor.lowering import register_lowering from torch._inductor.ir import TensorBox, ComputedBuffer @register_lowering(torch.ops.example.softmax_like) def lower_softmax_like(x, dim): # 这里可以进行 shape/dtype/stride 校验 # 如果条件不满足,可以返回 None 让 Inductor 回退到 eager if x.get_stride()[dim] != 1: return None # 返回一个描述循环融合的 IR 节点 # ...

我见过太多人卡在这一步:明明 lowering 注册了,逻辑也写了,但编译出来的代码就是不经过你的 kernel。最常见的两个原因,一个是自定义 op 的参数类型和你 lowering 函数签名对不上,另一个是返回的 IR 节点没有正确描述索引关系,导致 Inductor 认为输出无法索引。

请注意,register_loweringTritonTemplate都属于内部 API,不同 PyTorch 版本之间经常调整。你在生产代码里使用它们之前,务必先在当前版本里读一遍源码确认接口,不要照着旧博客硬抄。

3.4 用 IR 转储确认这次融合真的生效了

写完了不等于跑通了,一定要用工具把编译中间过程 dump 出来看一眼。PyTorch 2.x 里最省事的办法是打开 trace:

import torch._inductor.config as inductor_config inductor_config.trace.enabled = True

开启之后,Inductor 会往 debug 目录输出一整套文件,包括 FX 图变换前后的可读版本、Inductor IR 的中间状态、以及最终生成的 Triton/C++ 代码。你需要重点看的是这三个信息:

第一,变换后的 FX 图里还有没有你替换出来的那个自定义节点。如果新节点在后续某个 pass 里被 recognize 掉、重新拆回exp + sum + div,说明你的 lowering 没有生效,或者后续某个 decompose pass 不认识你的节点。第二,Inductor IR 的循环结构里,是不是真的只剩一个 kernel 的雏形,而不是又冒出来一堆中间 buffer。第三,最终 output_code 里的 kernel 数量。融合前是两个或者更多 kernel,融合正确的话应该明显变少。

这一步最能帮你区分"我的融合没跑"和"我的融合跑了但被后续 pass 拆回去了"这两种截然不同的失败模式。我自己的排查顺序通常是这样:先看output_code有没有自定义 kernel,没有就回头看 lowering;如果 lowering 没问题,再往前看 FX 变换后的图,检查替换是否成功。

4. 自定义 fusion 最容易踩的四个坑:缓存、别名、顺序和收益

4.1 编译缓存导致新 Pass “看起来没生效”

这是新手最容易误判的一个坑。你把自定义 pass 写好了,信心十足地运行脚本,结果跑出来的代码跟之前一模一样,时间也没有变少。你以为 pass 没跑,于是加了 print,结果 print 确实打出来了,代码还是没有变化。这个问题十有八九出在编译缓存上。

Inductor 默认会在磁盘上缓存编译产物,如果你的代码改动没有影响到缓存 key,比如你只改了一个模式匹配函数内部的逻辑,缓存 key 没变,那么这次 compile 会直接从缓存里加载旧结果,你的新 pass 根本没有被重新执行。这个现象特别隐蔽,因为你加了 print 也一样看不到,因为整个compile_fx流程在你开始打印前就已经被短路了。

处理方式有几种。最省事的是直接清缓存目录,通常是~/.cache/torchinductor或者当前项目下的.torchinductor;也可以用环境变量或者torch._inductor.config里的缓存开关(不同版本设置项有差异)临时禁用缓存。我的建议是:在开发和调试自定义 fusion 的阶段,就保持缓存禁用,直到代码稳定之后再打开,否则你会被缓存骗得团团转。

4.2 图形变叠加会让你的匹配被“毒化”

第二个大坑是 pass 之间的相互影响。你写的匹配逻辑是基于"理想图"设计的,但真实进入PostGradPassManager的图已经经过了前面多个 pass 的改造。比如你想匹配exp后面直接跟div,但事实上有可能是先有exp、然后插入了一个rand_like、再是div,因为某个正则化逻辑或者 dropout 插在中间。这种情况下你的 pattern 永远匹配不上。

更麻烦的是别名和 inplace 操作带来的"隐形边"。FX 图上call_function节点之间看起来是独立的小节点,但有些操作会共享存储。比如add_这类 inplace 算子,虽然 FX 里它是call_method,但因为它既读输入又改输入,如果你没有检查所有使用点就把它替换掉,后面原有图结构里的其他节点引用就会瞬间指到错误的数据上。

因此自定义 fusion 的匹配代码里,必须养成三个习惯:

  • 匹配节点时先看node.meta里的 shape、dtype、stride,发现不满足约束直接跳过;
  • 匹配一条链时记录每个节点的users数量,只处理使用计数完全符合预期的节点;
  • call_method,尤其是 inplace 版本,保持高度警惕,默认不匹配,除非你有明确的特殊处理。

4.3 单一顺序假设:一个 Pass 的产出是另一个 Pass 的输入

PostGradPassManager里的每个 pass 都会改变图的状态,这种改变会直接决定下一个 pass 能匹配到什么。这是图变换流水线里最本质也最容易被忽略的性质:pass 是有顺序的,顺序不是"可以随便排",而是"不同顺序可能得到完全不同的优化结果"。

举个例子,有些 decompose pass 会把某些高阶算子拆解成更底层的原子算子;如果你的 fusion pass 先跑,好不容易把softmax_like这种高阶节点融合出来了,后续 decompose pass 不认识它,又把它重新拆回原始算子序列,整个融合就白做了。反过来,如果你把 fusion pass 放在 decompose 之后跑,那它匹配的就是已经被拆碎的图,有些跨算子的结构信息可能已经找不回来。

这也是为什么我建议你在工程上做两件事。第一,在 pass 代码里输出足够明确的日志,至少记录你匹配了多少次、替换了多少个节点。第二,在把自定义 pass 接入流水线时,明确写清楚它依赖前面的哪个 pass 产生了什么形态的图,以及它不允许后面的哪些 pass 再碰它产生的节点。这两件事都能让问题在出现的第一时间被发现。

4.4 融合不是免费的:先量化收益再谈优化

最后一个坑比较反直觉:融合并不总是带来性能提升。很多点对点(pointwise)类融合确实稳赚不赔,因为减少了 kernel launch 和中间 buffer 的读写。但当你开始融合 reduction 和 pointwise 时,情况就会复杂很多。

一个典型的反面案例是:你把一个大的sum跟一个 pointwise 算子融合,结果导致 reduction 的并行度下降,Triton 里 tile 尺寸选不好,最后 kernel 的占用率反而不如两个独立 kernel。这种情况尤其容易出在输入形状不规则、最后一维不是连续维度、或者某些维度特别短的张量上。你融合之后生成的 kernel 可能在寄存器层面疯狂溢出,性能还不如不融合。

所以自定义 fusion 的真实工作流不是"写出来就完事",而是一定要在跑完之后做一次量化对比。对比的指标不光是 wall time,还要看 Profiler 里每个 kernel 的耗时、占用率、访存带宽。如果融合后的 kernel 耗时更高,先别急着否定 fusion 本身,试着调整 tile 或者换一种循环布局,依然没有起色的话就果断回退。我自己处理过好几个 case,融合逻辑没问题,但放到特定形状上就是负优化,这种时候保留原始 pattern 不融合反而是更正确的选择。

5. 一套可复现的自定义 fusion 调试验证流程

5.1 从日志和 dump 文件中读懂编译流水线

写自定义 fusion 最忌讳的是一遍遍地改代码、跑训练、看整体时间。整个过程太慢,噪音太大,而且无法定位问题出在哪个环节。正确做法是用 compile 的日志和 dump 文件做精细排查。

PyTorch 2.x 提供了多种观察手段。比较简单的是使用环境变量TORCH_LOGS,比如TORCH_LOGS="+dynamo,graph_code,output_code",这会输出 Dynamo 捕获后的 FX 图、每次图变换后的代码、以及最终生成的 GPU 代码。这个日志量会非常大,适合用来做深度排查,不适合日常开着跑。

日常开发我更推荐inductor_config.trace.enabled,它会在每次compile_fx调用时输出一套完整的、按阶段组织的 dump 文件。这里面最有用的是fx_graph_transformed.py(变换后的 FX 图)和output_code.py(生成的最终代码)。你可以直接读文件,然后对比其中是否有你的自定义节点和自定义 kernel。

把 dump 文件和自定义 pass 的命中日志配合起来看,你就有一张非常完整的流水线现场图。哪个阶段匹配成功、哪个阶段被改写、最终生成了什么样的代码,一目了然。

5.2 数值与性能的双重回归验证

排查完正确性之后,还有两项验证是自定义 fusion 上线前必须做的。第一项是数值一致性,第二项是性能收益。

数值验证不能只看整体 loss 是否收敛,而是要在小例子上直接对比原始算子和编译优化算子的输出。我的标准做法是:用固定随机种子生成一批输入,分别跑原始模型(eager 模式)和 torch.compile 优化后的模型,然后用torch.testing.assert_close对比输出和梯度,atolrtol通常设置在 1e-3 左右。如果是你自己写的 fusion,因为运算顺序发生了变化,数值不可能完全一致,但必须在小误差范围内。

性能验证方面,我建议写一个独立的小脚本,专门对比三个版本:纯 eager、torch.compile 默认优化、以及加入了自定义 fusion 后的优化。每个版本多跑几次取中位数,用torch.profiler记录 kernel 耗时,而不是只看脚本整体时间。脚本整体时间受到 Python 开销、图捕获开销、冷启动缓存等各种因素影响,干扰太大。

import torch from torch.profiler import profile, ProfilerActivity def run_eager(x): y = torch.exp(x) s = torch.sum(y, dim=-1, keepdim=True) return y / s def run_compiled(x): return torch.compile(run_eager)(x) x = torch.randn(4096, 1024, device="cuda") with profile(activities=[ProfilerActivity.CUDA]) as prof: for _ in range(10): run_eager(x) print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=10))

这个脚本虽然简单,却是所有 custom fusion 优化的起点。没有这套回归流程,你就无法区分"我的融合有没有生效"和"我的融合生效了但没有带来收益"这两个完全不同的结论。前一个是正确性问题,后一个是收益性问题,需要完全不同的后续处理方法。

最后分享一个实际经验:自定义 fusion 开始之前,先找一个最简单的 pointwise 融合跑通全链路,再逐步往 pattern 匹配里添加更多条件。我在给一个推荐系统模型的后处理节点做融合时,第一版也只匹配了三个节点,后面才逐渐扩展到带 mask、带 scale 的变体。图变换流水线这个东西,你理解得越细,越不容易做出"看起来很努力、实际没收益"的优化。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询