跳到主要内容

数据截至 (上游 commit c187ef3271d5)

06 · torch.compile 编译栈:从字节码到 Triton kernel

这一章讲什么: torch.compile(model) 这一行背后那套三段式流水线(Dynamo → AOTAutograd → Inductor)。读完你会知道:图是从哪抓的、抓不到怎么办、反向图是谁画的、Triton kernel 是谁写的。


1. 它要解决的小问题

eager 模式逐行调度算子,每个 *+ 都是一次 dispatcher 往返 + 一次 kernel 启动;小算子密集时,GPU 大部分时间在等 CPU 发活。

想要快,就得把一连串算子融合成少数几个大 kernel——这需要拿到「整段计算」的图。但 PyTorch 是 eager 框架,用户代码里可能混着 print、logging、数据依赖控制流,图没法靠 AST 静态分析拿到。

torch.compile 的回答:不读你的源码,直接在解释器层拦截字节码,跑一遍就知道你干了什么。


2. 流水线全景

你的代码 fn(x) / model(x)


① Dynamo 拦截 CPython 帧 → 逐条解释字节码 → FX 图 + guards
(torch/_dynamo/) 解释不了 → graph break,该段回落 eager


② AOTAutograd fake tensor 跑一遍前向 → 联合图 → partitioner 切出
(torch/_functorch/) 前向图与反向图(编译期就把反向也编译好)


③ Inductor 调度(哪些算子融合)→ codegen → Triton(GPU) / C++(CPU) kernel
(torch/_inductor/)


运行时 guards 命中 → 直接跑编译产物;不命中 → 重编译或回落

3. 第一段:Dynamo 截帧抓图

入口是 torch.compiletorch/__init__.py:3054),docstring 里把机制说得很直白(torch/__init__.py:3074-3080):

for every frame executed within the compiled region, we will attempt to compile it and cache the compiled result on the code object for future use. A single frame may be compiled multiple times if previous compiled results are not applicable for subsequent calls (this is called a "guard failure").

机制三层:

  • 帧拦截torch._dynamo.optimizetorch/_dynamo/eval_frame.py:1786)给目标函数的 code object 装上帧求值回调(convert_frame.catch_errors_wrappertorch/_dynamo/eval_frame.py:357);CPython 每次执行这个帧,先问 Dynamo。
  • 字节码解释InstructionTranslatortorch/_dynamo/symbolic_convert.py:5587)是一个用 Python 写的 Python 字节码解释器,逐条 steptorch/_dynamo/symbolic_convert.py:1703)执行指令;遇到 PyTorch 算子调用就记进 FX 图,遇到普通 Python(如算术、list 操作)就直接求值成常量。
  • guards:抓图时同时记下前提(张量形状/dtype/设备、模块属性身份、全局变量值……)。下次执行先查 guards,命中就走缓存的编译产物。

graph break 是安全阀:遇到解释不了的指令(如 print、某些 C 扩展),Dynamo 把已抓的图切出来编译,剩余部分继续用普通解释执行。所以 fullgraph=False(默认)时编译几乎不会「失败」,只会「收益打折」。


4. 第二段:AOTAutograd 连反向一起抓

eager 的反向图是运行时才挂的(第 3 章);编译栈想优化反向,就得提前把它画出来。这就是 AOTAutograd 的活。

入口 aot_functiontorch/_functorch/aot_autograd.py:712),docstring 把三步说清了(torch/_functorch/aot_autograd.py:728-737):

traces the forward and backward graph ahead of time, and generates a joint forward and backward graph. partition_fn is then used to separate out forward and backward graphs.

做法:

  1. fake tensor 探测:用只有元数据(形状/dtype/device)没有数据的 fake tensor 跑一遍前向,经过 dispatcher 时算子被记成图节点——这一步复用了第 2 章的 dispatch 机制(挂上 FakeTensor 对应的 dispatch 层)。
  2. 联合图:把前向和反向描成一张联合图。
  3. partitioner 切图:切成前向图和反向图;这一步还能做 recomputation(激活重算)之类的显存/时间权衡。

产出交给编译后端:默认后端 Inductor 通过 compile_fxtorch/_inductor/compile_fx.py:3084)接住两张图。后端是可替换的——注册表在 torch/_dynamo/backends/registry.py_BACKENDS:81register_backend:87),这就是 TensorRT、OpenXLA 等第三方后端插入的位置。


5. 第三段:Inductor 出 kernel

Inductor 拿到 FX 图后做两件事:

  • 调度与融合Schedulertorch/_inductor/scheduler.py:4922)决定哪些节点能融进同一个 kernel——逐元素链、pointwise+reduction 是典型融合对(FusedSchedulerNode.fusetorch/_inductor/scheduler.py:3358)。
  • 代码生成:GPU 上由 TritonKerneltorch/_inductor/codegen/triton.py:3301)把每个融合组写成一段 Triton kernel 源码codegen_kerneltorch/_inductor/codegen/triton.py:7370),编译加载执行;CPU 走 C++ codegen。

效果:一段 y = (a * b).relu().sum(dim=1) 在 eager 下是 3+ 次 kernel 启动、中间结果反复读写显存;编译后通常是一个 kernel,一次读写。


6. 原理演示:三段各看见什么

# 示意,非源码
@torch.compile # ① 给 fn 的 code object 装帧拦截
def fn(x, w):
print("compiling…") # graph break:这句留在 eager
return (x @ w).relu().sum() # 这段被抓成 FX 图

# ② AOTAutograd 看到的(伪 FX):
# joint graph: mm → relu → sum,以及它们各自的 backward 公式
# ③ Inductor 把 relu+sum 融成一个 Triton kernel;mm 可走 cuBLAS/模板

重点看:编译栈里的反向不再是「运行时挂图」,而是「编译期算好的另一张图」——loss.backward() 跑的是 Inductor 编出来的反向 kernel。


7. 坑与边界

  • guard 失败会静默重编译,代码里形状/控制流频繁变化时编译缓存抖动,可能比 eager 还慢;有 recompile_limit 兜底。
  • graph break 不报错(默认 fullgraph=False),性能不达预期时第一件事是用 TORCH_LOGS=graph_breaks 之类查断在哪。
  • 动态形状是 guards 的常客:shape 变化触发特化重编译,dynamic=True 走符号形状路径,但覆盖面和收益都要实测。
  • 训练时反向也编译意味着第一次迭代很慢(warmup),短任务不划算。
  • 不是所有后端路径等价:Inductor 对数值精度有自己的取舍(如融合的 reduction 顺序),对精度敏感场景要开严格模式对比。
  • 看不出来:Inductor 模板调度(matmul 用 Triton 模板 vs cuBLAS 的选择逻辑)的全部启发式——在 torch/_inductor/ 的 scheduler 与 template 代码里,细节建议按需深入。

8. 本章代码地图

主题文件路径符号名
编译入口torch/__init__.pytorch.compile
帧拦截torch/_dynamo/eval_frame.pyoptimizecatch_errors_wrapper
字节码解释器torch/_dynamo/symbolic_convert.pyInstructionTranslator.step
帧编译torch/_dynamo/convert_frame.pyconvert_frame_compile
AOT 抓图torch/_functorch/aot_autograd.pyaot_function
后端注册表torch/_dynamo/backends/registry.pyregister_backend_BACKENDS
Inductor 入口torch/_inductor/compile_fx.pycompile_fx
融合调度torch/_inductor/scheduler.pySchedulerFusedSchedulerNode.fuse
Triton 代码生成torch/_inductor/codegen/triton.pyTritonKernelcodegen_kernel