首页
/ PyTorch 实验性编译 API 指南:基于 functorch.compile 的 AOT Autograd 使用与原理

PyTorch 实验性编译 API 指南:基于 functorch.compile 的 AOT Autograd 使用与原理

2026-09-09 20:23:35作者:齐添朝

AOT Autograd(Ahead-of-Time Autograd)是 PyTorch 中一项实验性的提前编译功能:它允许在训练前提前捕获前向与反向计算图,并在此基础上轻松接入各类编译器(TorchScript、NNC、NVFuser 等),从而为加速 PyTorch 模型训练提供一个易于修改的纯 Python 开发环境。本指南以 functorch/docs/source/aot_autograd.rst 文档为主体,结合仓库中 functorch.compile 命名空间下的真实实现(torch/_functorch/aot_autograd.pytorch/_functorch/compilers.pytorch/_functorch/partitioners.py)与示例代码,读完你将掌握 aot_function / aot_module 的调用方式、编译器的编写范式、切分器(partitioner)的作用,以及如何用 memory_efficient_fusion 做显存高效融合。

⚠️ 实验性警告:AOT Autograd 属于实验性特性,API 很可能发生变更。官方正在寻求社区反馈——如果你有意使用 AOT Autograd 并在使用中需要帮助或有建议,欢迎通过 GitHub Issue 反馈(对应源码注释与文档中均明确标注了这一点)。


1. AOT Autograd 是什么:核心概念与设计动机

functorch/docs/source/aot_autograd.rst 的定义来看,AOT Autograd 的核心能力可以概括为三点:

  1. 提前捕获(ahead of time capture):在真正的训练迭代开始之前,就把一个 Python 函数(或 nn.Module)的前向计算与反向计算完整地捕获为图结构;
  2. 易于集成编译器(easy integration with compilers):捕获得到的图是标准 torch.fx.GraphModule,可以直接交给任意编译器处理;
  3. 纯 Python 可 hack 的开发环境:由于整个捕获与编译流程都用 Python 编写,研究者可以快速修改图、插入自定义 pass,再也不用依赖黑盒式的编译器内部机制。

AOT Autograd 当前位于 functorch.compile 命名空间。在仓库中,该命名空间由 functorch/compile/init.py 对外暴露,其内部实现统一来自 torch._functorch 下的三个核心模块:

1.1 为什么需要"提前"捕获反向图

常规的 PyTorch 训练通过 autograd 引擎在反向传播时动态构建计算图,编译器难以在运行前看到完整的反向结构。AOT Autograd 的思路是:先完整追踪一次前向与反向(生成 joint graph,即前向+反向联合图),再交给 partitioner 切分为独立的前向图和反向图,最后分别用 fw_compilerbw_compiler 编译。这样编译器拿到的就是结构完整、可整体优化的图,而不是逐算子调度的 eager 执行。


2. 编译 API 一览(experimental)

原文档将 functorch.compile 的公开 API 分为三类,仓库中的实际导出与之一一对应(见 functorch/compile/init.py):

类别 API 仓库实现位置
编译入口 aot_function torch/_functorch/aot_autograd.py#L712
编译入口 aot_module torch/_functorch/aot_autograd.py#L847
编译入口 memory_efficient_fusion torch/_functorch/compilers.py#L278
切分器 default_partition torch/_functorch/partitioners.py#L1617
切分器 min_cut_rematerialization_partition torch/_functorch/partitioners.py#L4164
编译器 nop torch/_functorch/compilers.py#L134
编译器 ts_compile torch/_functorch/compilers.py#L68

2.1 aot_function:追踪并编译任意 Python 函数

aot_function 接收一个返回一个或多个 Tensor 的 Python 函数,通过 torch dispatch 机制提前追踪其前向与反向图,并将生成的两个图分别交给 fw_compilerbw_compiler 编译。其完整签名如下(见 aot_autograd.py#L712-L727):

def aot_function(
    fn: Callable[_P, _R],
    fw_compiler: AOTDispatchCompiler,
    bw_compiler: AOTDispatchCompiler | None = None,
    partition_fn: Callable[..., Any] = default_partition,
    decompositions: dict[OpOverload, Callable[..., Any]] | None = None,
    num_params_buffers: int = 0,
    keep_inference_input_mutations: bool = False,
    inference_compiler: AOTDispatchCompiler | None = None,
    *,
    dynamic: bool = False,
    enable_log: bool = True,
    disable_functionalization: bool = False,
    _disable_torch_fn_metadata_mode: bool = False,
) -> Callable[_P, Any]:

各参数语义(依据 aot_autograd.py#L743-L777 的 docstring):

  • fn:待编译的 Python 函数,参数为一个或多个,返回值必须是一个或多个 Tensor;
  • fw_compiler:接收一个包含 ATen 算子的 FX 图与输入参数,返回一个语义等价的 Callable;
  • bw_compiler:同上,但针对反向图;默认为 None,此时回退到 fw_compiler
  • partition_fn:接收 joint(前向+反向联合)图,将其切分为独立的前向图与反向图。切分器可以在此处实现重计算(recomputation)等优化
  • decompositions:将较大的 ATen 算子分解为核心/更简单算子的字典,便于后端编译器支持;
  • inference_compiler:当不需要 autograd(推理场景)时被调用;默认为 None,回退到 fw_compiler
  • dynamic:是否以动态形状(dynamic shapes)进行追踪;
  • keep_inference_input_mutations:推理时是否保留对输入的就地修改语义。

2.2 aot_module:编译 nn.Module

aot_module 是面向 nn.Module 的便捷封装:它将模块的参数与 buffer 提升(lift)为输入,构造一个新的可调用对象,再经由 aot_function 完成追踪与编译(见 aot_autograd.py#L847-L879)。其参数与 aot_function 保持一致,返回值是保留了原模块 eager 行为、但前向/反向图已被编译的 nn.Module

2.3 memory_efficient_fusion:显存高效融合

memory_efficient_fusion 是文档中列出的第三个编译入口,实现在 torch/_functorch/compilers.py#L278-L314。它的内部工作方式非常清晰:

  1. 使用 min_cut_rematerialization_partition 切分器,对 joint graph 做最小割重计算(等价于自动化的 activation checkpointing 式重计算),以显存换取计算;
  2. 使用 ts_compile(TorchScript 编译 + 冻结)编译前向与反向图;
  3. 默认配置为:
config = {
    "fw_compiler": ts_compile,
    "bw_compiler": ts_compile,
    "partition_fn": min_cut_rematerialization_partition,
    "decompositions": default_decompositions,
}
config.update(kwargs)   # 允许调用方覆盖任意配置

如果传入的是 nn.Module 则走 aot_module,否则走 aot_functiondefault_decompositions 定义在同文件 compilers.py#L242-L269,包含 gelu_backwardleaky_relu_backwardsigmoid_backwardtanh_backwardsilu_backwardelu_backwardcudnn_batch_normmasked_fill 等常见算子的分解规则。


3. Partitioners(切分器):joint graph 的拆分与重计算

AOT Autograd 追踪到的是前向+反向联合图,必须经过 partitioner 切分后才能分别编译。文档列出了两个切分器:

3.1 default_partition:朴素切分

default_partition 位于 torch/_functorch/partitioners.py#L1617,实现朴素的前向/反向切分:前向图中需要保留给反向使用的中间激活(activations)直接保存在前向输出中,不做任何重计算优化。适合作为正确性基线使用。

3.2 min_cut_rematerialization_partition:最小割重计算

min_cut_rematerialization_partition 位于 torch/_functorch/partitioners.py#L4164,是显存优化的核心。它把"哪些激活该保存、哪些该在反向时重新计算"建模为图上的最小割(min-cut)问题,通过权衡保存激活的显存成本与重计算的计算成本,自动选择一组重计算节点。这与 PyTorch 官方 dev-discuss 中提出的 "Min-cut optimal recomputation (i.e. activation checkpointing) with AOTAutograd" 方案一脉相承(仓库 functorch/COMPILE_README.md 中给出了该讨论的链接)。

3.3 draw_graph:可视化辅助

切分器模块还提供了 draw_graphpartitioners.py#L4425),用于将 FX 图导出为 svg 可视化文件,方便调试图结构。


4. Compilers(编译器):从图到可执行代码

4.1 nop:什么都不做的编译器

nop 位于 torch/_functorch/compilers.py#L134-L143,直接原样返回输入的 FX GraphModule。它的用途是正确性对照(accuracy check):在接入真实编译器之前,先用 nop 验证 AOT Autograd 追踪出的图本身语义正确。

4.2 ts_compile:TorchScript 编译管线

ts_compile 位于 torch/_functorch/compilers.py#L68-L116,是一条完整的 TorchScript 编译管线,按顺序执行:

  1. strip_overloads(fx_g):去除 ATen 算子的 overload 后缀;
  2. _to_copy 节点规范化到 aten.to,并把 torch.device kwargs 转为字符串(TorchScript 兼容性处理);
  3. fx_g.graph.lint() + recompile():校验并重编译图;
  4. torch.jit.script(fx_g):将 FX 图脚本化为 TorchScript;
  5. torch._C._jit_pass_remove_mutation:移除可变性;
  6. torch.jit.freeze(f.eval()):冻结推理模式下的子图——注意文档与源码都强调,即使该图用于训练,这里的 eval() 也是安全的(见 functorch/COMPILE_README.md 中的示例注释);
  7. torch.jit.optimize_for_inference(f):推理优化;
  8. 若输入不是 fake tensor,则用真实输入执行一次以触发编译/验证。

同文件还提供了 simple_ts_compile(仅 script + freeze)与 nnc_jitaot_function(f, simple_ts_compile) 的简写,见 compilers.py#L238-L239),后者被 functorch/examples/compilation/simple_function.py 中的基准示例使用。


5. 实战示例:打印、可视化与 TorchScript 编译

以下示例完整取自仓库 functorch/COMPILE_README.md(该文件与 rst 文档是同一主题的配套 README),是 AOT Autograd 最直接的入门代码。

5.1 打印前向与反向图

from functorch.compile import aot_function, aot_module, draw_graph
import torch.fx as fx
import torch

# 该函数打印 FX 图(前向/反向)
def print_graph(name):
    def f(fx_g: fx.GraphModule, inps):
        print(name)
        print(fx_g.code)
        return fx_g
    return f

def f(x):
    return x.cos().cos()

nf = aot_function(f, fw_compiler=print_graph("forward"), bw_compiler=print_graph("backward"))
nf(torch.randn(3, requires_grad=True))

注意:nf 编译后的函数仍然可以通过 autograd 继续反向传播——你可以在调用前后随意插入其他算子,例如:

inp = torch.randn(3, requires_grad=True)
inp = inp.cos()
out = nf(inp)
out = out.sin().sum().backward()   # 依然可以反传

5.2 将图导出为 SVG 可视化

def graph_drawer(name):
    def f(fx_g: fx.GraphModule, inps):
        draw_graph(fx_g, name)     # 生成名为 name 的 svg 文件
        return fx_g
    return f

aot_function(f, fw_compiler=graph_drawer("forward"), bw_compiler=graph_drawer("backward"))(
    torch.randn(3, requires_grad=True)
)

5.3 对 nn.Module 使用 aot_module

from torchvision.models import resnet18
aot_module(resnet18(), print_graph("forward"), print_graph("backward"))(
    torch.randn(1, 3, 200, 200)
)
# 输出较长,此处省略

aot_module 会把 resnet18 的参数和 buffer 提升为图输入(对应 aot_autograd.py#L872-L879 中的 functional_call),因此生成的图是"无状态"的、便于编译器整体优化。

5.4 集成 TorchScript:一个完整的编译管线

def f(x):
    return x.cos().cos()

def ts_compiler(fx_g: fx.GraphModule, inps):
    f = torch.jit.script(fx_g)
    print(f.graph)
    f = torch.jit.freeze(f.eval())   # 用于训练同样安全
    return f

aot_function(f, ts_compiler, ts_compiler)(torch.randn(3, requires_grad=True))

这段代码即仓库内置 ts_compile 的手写简化版。在实际项目中,建议直接使用 torch/_functorch/compilers.py 提供的 ts_compile,因为它额外处理了 overload 去除、设备参数归一化、mutation 消除与推理优化等细节。

5.5 更多示例与正确性校验

仓库 functorch/examples/compilation/ 下还提供了配套示例:

⚠️ 注意:该目录下的 README.md 明确警告:编译功能目前非常实验性,示例可能无法开箱即用。

此外,若要快速验证追踪出的图语义是否正确,可用 nop 编译器与原 eager 函数对比输出;torch/_functorch/compilers.py#L218-L227 还提供了 debug_nop——它返回一个 DebugInterpreter,在解释执行 FX 图的同时校验每个节点的 dtype、shape 与 stride 是否与追踪时的元数据一致,是排查图捕获问题的利器。


6. 底层原理:从 torch dispatch 到缓存复用

aot_function 的内部流程(见 aot_autograd.py#L799-L844)可以概括为:

  1. 展平输入:通过 pytree.arg_tree_leaves 将嵌套的 args/kwargs 展平为 flat args;
  2. 构造 fake mode 与 shape envconstruct_fake_mode 使用 fake tensor(torch/_subclasses/fake_tensor.py)在不真正执行算子的前提下模拟张量元数据,使追踪成为可能;
  3. 一次追踪、两次编译create_aot_state + aot_stage1_graph_capture 捕获 joint graph,aot_stage2_compile 调用 partition_fn 切分,并分别交给 fw_compilerbw_compilerinference_compiler 编译;
  4. 结果缓存:编译结果保存在闭包变量 cached_res 中,后续每次调用直接复用已编译的函数,避免重复追踪——这意味着首次调用承担全部编译开销,之后的调用走编译后的代码路径。

nn.Module 而言,aot_module 的核心技巧是参数/buffer 提升:它构造一个 functional_call,把命名参数与 buffer 合并后通过 torch.func.functional_call 注入模块前向,从而让整个前向变成"输入即参数"的纯函数,这正是 aot_autograd.py#L869-L879functional_call 的实现,也解释了为什么图中不包含模块状态。


7. 总结与使用建议

围绕 functorch/docs/source/aot_autograd.rst 的骨架,我们可以提炼出 AOT Autograd 的完整使用路径:

  1. 选择入口:函数用 aot_function,模块用 aot_module,需要显存优化直接上 memory_efficient_fusion
  2. 选择编译器:验证正确性用 nop / debug_nop,调试看结构用 print_compile / draw_graph_compile,落地加速用 ts_compile / nnc_jit
  3. 选择切分器:默认 default_partition;追求显存效率时换成 min_cut_rematerialization_partition,让 AOT Autograd 自动决定哪些激活需要重计算;
  4. 注意实验性边界:文档与源码多处标注 "experimental and likely to change",API 签名(如新增的 dynamicdisable_functionalization 参数)都可能在后续版本调整,生产环境落地前应锁定版本并关注变更。

从更宏观的视角看,AOT Autograd 捕获的"前向+反向联合图 + 可插拔 partitioner/compiler"架构,正是 PyTorch 编译栈(如 torch._dynamotorch._inductor)的基础组件之一——理解这套 API,也就理解了 PyTorch 训练编译流水线的核心抽象。

登录后查看全文
热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
docsdocs
暂无描述
Markdown
899
5.83 K
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.76 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.35 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
925
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.84 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
601
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
395
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.04 K
525