
跑过几轮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 执行关系非常紧密里面可能还残留着getattr、assert、Python 控制流留下的子图调用。它的职责是把你写的 Python 代码翻译成算子调用序列但远没有到适合做优化的程度。比如你写一个循环Dynamo 可能捕获成多个call_function节点也可能整体变成一个 subgraph里面继续走 Python 解释器这在后续优化里都非常烫手。第二层是AOTAutograd 产出的、已经展开自动求导的 Graph。这一层的图里不再有autograd.Function那样隐含的求导逻辑而是显式地把 forward 和 backward 需要的算子都展开成一张大图。也正因为求导被展开了很多原来被隐藏在 autograd 黑盒里的冗余计算、可以被消掉的中间变量现在都暴露出来了。这个阶段产出的图才真正适合做代数化简、公共子表达式消除、死代码消除等操作。第三层是Inductor IR。FX Graph 仍然是算子级别的抽象而 Inductor 会把算子再拆成更细粒度的循环、buffer、布局描述。它不再是图的形态而是一组Pointwise、Reduction、TemplateBuffer等 IR 节点。到了这一层PassManager 的职责开始转向循环级变换、调度顺序、缓存局部性这类底层问题。所以当大家说图变换流水线的时候一定要先确认自己站在哪一层。每一层的图结构不同数据依赖的表达方式不同能做和不能做的变换也完全不同。1.2 为什么必须多出一个 PostGrad 阶段不同编译器对图变换流水线的划分很不一样。LLVM 里是ModulePassManager加FunctionPassManagerMLIR 里有从Module到Func再到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 做的事情大致可以分成四类。一类是移除无意义节点。自动求导展开后经常产生很多view、clone、detach、alias这类节点它们的存在往往只是为了让 autograd 的链式求导规则能够成立。在图变换阶段它们已经没有任何价值留下来只会干扰后续 pattern 匹配。因此会有一批专门做 noop 消除的 pass 把它们删掉。第二类是代数化简与维度整理。比如把split之后再cat这种逆操作消除把多次unsqueeze/squeeze合并成一次把permute链化简成单个permute。这些操作在展开后的图上特别常见因为反向传播里为了对齐梯度 shape会大量插入 reshape 类算子。第三类是pattern 识别式融合。像fuse_attention、fuse_conv_bn、fuse_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, keepdimTrue)这种 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, keepdimTrue) - 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 上做子图同构查找。真正实战时还需要补充很多细节比如必须验证sum的dim参数是尾部维度、keepdimTruediv_node的第二个参数必须是那个 sum 输出而不是别的标量还要检查这些节点的使用计数确保y_node除了被sum和div使用之外没有被其他地方引用。匹配到之后替换逻辑并不复杂但有一个容易忽略的坑不能直接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_lowering和TritonTemplate都属于内部 API不同 PyTorch 版本之间经常调整。你在生产代码里使用它们之前务必先在当前版本里读一遍源码确认接口不要照着旧博客硬抄。3.4 用 IR 转储确认这次融合真的生效了写完了不等于跑通了一定要用工具把编译中间过程 dump 出来看一眼。PyTorch 2.x 里最省事的办法是打开 traceimport 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_LOGSdynamo,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对比输出和梯度atol和rtol通常设置在 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, keepdimTrue) return y / s def run_compiled(x): return torch.compile(run_eager)(x) x torch.randn(4096, 1024, devicecuda) with profile(activities[ProfilerActivity.CUDA]) as prof: for _ in range(10): run_eager(x) print(prof.key_averages().table(sort_bycuda_time_total, row_limit10))这个脚本虽然简单却是所有 custom fusion 优化的起点。没有这套回归流程你就无法区分我的融合有没有生效和我的融合生效了但没有带来收益这两个完全不同的结论。前一个是正确性问题后一个是收益性问题需要完全不同的后续处理方法。最后分享一个实际经验自定义 fusion 开始之前先找一个最简单的 pointwise 融合跑通全链路再逐步往 pattern 匹配里添加更多条件。我在给一个推荐系统模型的后处理节点做融合时第一版也只匹配了三个节点后面才逐渐扩展到带 mask、带 scale 的变体。图变换流水线这个东西你理解得越细越不容易做出看起来很努力、实际没收益的优化。