PyTorch torch.compile 入门实战指南:从三角函数融合到预训练模型加速
本文是 PyTorch 官方 Torch Compiler 用户指南的入门章节(对应仓库文档 torch.compiler_get_started.md)的深度展开。文章围绕 torch.compile 的推理用法展开:先用一个三角函数点运算示例建立对内核融合(fusion)的直觉,再通过调试环境变量查看 TorchInductor 生成的实际 Triton 内核,最后以 ResNet50、HuggingFace BERT 与 TIMM 预训练模型演示"一行代码"接入编译加速的完整路径。读完本文,你将掌握 torch.compile(backend="inductor") 的基本调用方式、后端选择方法、调试与代码生成检查手段,并能在 CPU / NVIDIA GPU / Intel GPU 上复现全部示例。建议先阅读总览章节 torch.compiler_dynamo_overview.md(即文档中的 torch.compiler_overview 引用目标)以建立整体概念。
第一个 torch.compile 示例:从点运算理解内核融合
TorchDynamo 与 TorchInductor 的设计目标之一,是让 torch.compile 对用户编写的模型"开箱即用"。我们先看一个最简单的推理示例——它演示了 torch.cos() 与 torch.sin() 这两个典型的逐元素(pointwise)算子:
import torch
def fn(x):
a = torch.cos(x)
b = torch.sin(a)
return b
new_fn = torch.compile(fn, backend="inductor")
input_tensor = torch.randn(10000).to(device="cuda:0")
a = new_fn(input_tensor)
::: note
运行该脚本需要机器上至少有一块 GPU。如果没有 GPU,可以删掉代码中的 .to(device="cuda:0"),脚本就会在 CPU 上运行;也可以把设备改为 xpu:0,从而在 Intel® GPU 上运行。
:::
这个例子可能不会带来显著的性能提升,但它能帮你建立对 torch.compile 工作方式最直观的认识:你只需要把想要优化的函数或模型用 torch.compile 包一层,然后像调用普通函数一样调用编译后的对象即可。从仓库源码看,torch.compile 在 torch/compiler/init.py 中被统一暴露,底层由 TorchDynamo 负责图捕获,再交由指定的后端(默认与最常用的是 inductor)完成代码生成与优化。
torch.relu() 是另一个更广为人知的点运算算子。在 eager 模式下,点运算的效率是不理想的:每一个点运算算子都需要从内存读取一个张量、做一点修改、再把结果写回内存。当多个点运算串联时,这种"读-写-读-写"的往返会成倍放大内存流量。
TorchInductor 执行的最重要的优化就是融合(fusion)。以上面的例子为例,eager 模式需要 2 次读取(x、a)和 2 次写入(a、b),而融合后只需要 1 次读取(x)和 1 次写入(b)。这对于新一代 GPU 尤其关键,因为此时性能瓶颈已经从"计算能力"(GPU 每秒能完成的浮点运算量)转移到了内存带宽(数据送达 GPU 的速度)。
TorchInductor 的另一大优化:CUDA Graphs
除内核融合外,TorchInductor 提供的另一项重要优化是对 CUDA graphs 的自动支持。
CUDA graphs 通过一次性捕获一段 GPU 操作序列并整体重放,消除了从 Python 程序逐个启动 kernel 所带来的 CPU 开销(即 launch overhead)。对于现代 GPU 上那些执行时间极短、kernel 数量众多的计算图来说,这种启动开销往往占据主导地位,因此 CUDA graphs 的意义十分突出。
仓库中的 torch/_dynamo/backends/cudagraphs.py 完整实现了 CUDA graphs 后端支持,模块文档明确说明其职责包括:为前向与反向传播创建和管理 CUDA graph、检测并处理输入突变、进行设备兼容性检查、与 TorchInductor 的 cudagraph trees 集成等。它支持两种主要模式:
cudagraphs:完整的 CUDA graph 支持,同时优化前向与反向传播;cudagraphs_inner:用于基准测试的低层 CUDA graph 实现。
这也是官方文档在介绍完 inductor 后,建议读者用 cudagraphs 后端做下一个尝试的原因。
用 TORCH_COMPILE_DEBUG 查看生成的内核代码
TorchDynamo 支持多种后端,而 TorchInductor 的工作方式是生成 Triton 内核。把上面的示例保存为 example.py,然后运行:
TORCH_COMPILE_DEBUG=1 python example.py
脚本执行过程中,终端会打印出 DEBUG 日志。在日志接近末尾的位置,你会看到一个指向某个文件夹的路径,其中包含 torchinductor_<你的用户名> 目录。在这个目录中可以找到 output_code.py 文件,里面就是生成的 kernel 代码,类似下面这样:
@pointwise(size_hints=[16384], filename=__file__, triton_meta={'signature': {'in_ptr0': '*fp32', 'out_ptr0': '*fp32', 'xnumel': 'i32'}, 'device': 0, 'constants': {}, 'mutated_arg_names': [], 'configs': [AttrsDescriptor(divisible_by_16=(0, 1, 2), equal_to_1=())]})
@triton.jit
def triton_(in_ptr0, out_ptr0, xnumel, XBLOCK : tl.constexpr):
xnumel = 10000
xoffset = tl.program_id(0) * XBLOCK
xindex = xoffset + tl.arange(0, XBLOCK)[:]
xmask = xindex < xnumel
x0 = xindex
tmp0 = tl.load(in_ptr0 + (x0), xmask, other=0.0)
tmp1 = tl.cos(tmp0)
tmp2 = tl.sin(tmp1)
tl.store(out_ptr0 + (x0 + tl.zeros([XBLOCK], tl.int32)), tmp2, xmask)
::: note 以上代码片段只是一个示例。根据硬件不同,你实际看到的生成代码可能有所差异。 :::
你可以借此验证 cos 和 sin 的融合确实发生了:cos 与 sin 位于同一个 Triton kernel 内部,中间的临时变量保存在访问速度极快的寄存器中,而不是被写回显存。整段内核代码是 Python 写的,即便你从未编写过太多 CUDA kernel,也比较容易读懂。
从仓库源码看,TORCH_COMPILE_DEBUG 是 TorchInductor 调试机制的"总开关"。在 torch/_inductor/config.py 中可以看到:
# master switch for all debugging flags below
enabled = os.environ.get("TORCH_COMPILE_DEBUG", "0") == "1"
除总开关外,还有一系列配套环境变量,例如:
TORCH_COMPILE_DEBUG_SAVE_REAL:保存真实张量,便于排查数值问题;TORCH_COMPILE_DEBUG_EXTEND:把 Inductor kernel 的栈追踪信息回填到 PyTorch profiler 时间线中;TORCH_COMPILE_DEBUG_MAX_EVENTS:profiler 时间线后处理的最大 trace 事件数(默认 500000),超出会跳过溯源处理以避免内存溢出;INDUCTOR_PROVENANCE:控制溯源跟踪级别(1 为正常,2 为 basic)。
调试目录的实际创建逻辑位于 torch/_inductor/debug.py 的 DebugContext.create_debug_dir(),它会基于 config.trace.debug_dir 或系统临时目录生成 torchinductor 子目录。想深入学习 Triton 的性能特性,可以阅读 Triton 官方文档(triton-lang.org)。
选择后端:torch.compiler.list_backends()
inductor 并不是唯一可用的后端。你可以在 Python REPL 中运行 torch.compiler.list_backends() 查看当前环境中所有可用的后端名称,然后尝试其中的 cudagraphs。
从源码看,后端发现与注册机制位于 torch/_dynamo/backends/registry.py:
register_backend(compiler_fn, name=None, tags=()):把编译函数注册到后端表中,register_debug_backend与register_experimental_backend分别是带debug、experimental标签的便捷注册方式;lookup_backend(compiler_fn):把后端字符串展开为对应的编译函数;如果名称不存在,会借助difflib.get_close_matches给出相近名称的纠错建议,并抛出InvalidBackend;list_backends(exclude_tags=("debug", "experimental")):返回所有可传入torch.compile(..., backend="name")的合法字符串,默认排除debug与experimental标签的后端;_lazy_import()会加载torch._dynamo.backends下所有子模块,并通过entry_points(分组名torch_dynamo_backends)发现第三方注册的后端。
对应地,公开 API torch.compiler.list_backends 在 torch/compiler/__init__.py 中做了转发封装。仓库自带的动态后端模块位于 torch/_dynamo/backends 目录,包括 inductor.py、cudagraphs.py、onnxrt.py、tensorrt.py、tvm.py、torchxla.py、distributed.py 以及 debugging.py 等。第三方库也可以通过声明 torch_dynamo_backends entry point 贡献自己的后端——这就是 list_backends() 的输出会随环境而变化的原因。
上手真实模型:ResNet50
接下来用一个真实模型——来自 PyTorch Hub 的 ResNet50——来体验 torch.compile:
import torch
model = torch.hub.load('pytorch/vision:v0.10.0', 'resnet50', pretrained=True)
opt_model = torch.compile(model, backend="inductor")
opt_model(torch.randn(1, 3, 64, 64))
torch.hub.load 会从 pytorch/vision 仓库拉取模型定义与预训练权重,torch.compile 在首次调用时完成图捕获与代码生成,后续调用即可复用编译产物。
使用预训练模型:HuggingFace Transformers 与 TIMM
PyTorch 用户经常使用来自 transformers 或 TIMM 的预训练模型,而 TorchDynamo 与 TorchInductor 的设计目标之一,就是与人们想要编写的任何模型开箱即用地配合工作。
优化 HuggingFace BERT
下面直接从 HuggingFace Hub 下载一个预训练 BERT 模型并优化它:
import torch
from transformers import BertTokenizer, BertModel
tokenizer = BertTokenizer.from_pretrained('bert-base-uncased')
model = BertModel.from_pretrained("bert-base-uncased").to(device="cuda:0")
model = torch.compile(model, backend="inductor") # 这是唯一改动的一行代码
text = "Replace me by any text you'd like."
encoded_input = tokenizer(text, return_tensors='pt').to(device="cuda:0")
output = model(**encoded_input)
这里值得强调:对于既有模型,接入编译优化通常只需要修改一行代码——把 model 用 torch.compile(model, backend="inductor") 包起来,其余推理代码完全不用动。这正是 TorchDynamo 图捕获能力的体现:它能在 Python 层透明地捕获模型前向计算,再交给 Inductor 做后端代码生成。
如果你把 model 和 encoded_input 上的 to(device="cuda:0") 都删掉,那么 Triton 将改而生成经过优化、面向 CPU 运行的 C++ 内核。你可以同样去查看 BERT 对应的 Triton 或 C++ 内核——它们比前面的三角函数示例复杂得多,但同样可以快速浏览,体会 PyTorch 的计算图是如何被翻译为底层 kernel 的。
优化 TIMM 模型
再来看一个 TIMM 的例子:
import timm
import torch
model = timm.create_model('resnext101_32x8d', pretrained=True, num_classes=2)
opt_model = torch.compile(model, backend="inductor")
opt_model(torch.randn(64, 3, 7, 7))
注意这里输入形状是 (64, 3, 7, 7),与通常的 ImageNet 分辨率不同——torch.compile 并不要求输入必须与预训练时的形状完全一致,它会基于实际运行时的张量形状与元数据进行编译与优化(动态形状与重编译的细节可进一步阅读仓库中的 torch.compiler_dynamic_shapes.md)。
下一步学习路径
本节通过几个推理示例,帮助你建立了对 torch.compile 工作方式的基本理解。接下来可以按以下路径继续深入:
- 训练场景:阅读 PyTorch 官方《torch.compile tutorial on training》教程,了解编译对训练循环的加速;
- API 参考:查看 torch.compiler_api.md(即文档中的
torch.compiler_api引用目标),掌握torch.compile的完整参数(如mode、fullgraph、dynamic、options等); - 细粒度追踪:阅读 torch.compiler_fine_grain_apis.md(即文档中的
torchdynamo_fine_grain_tracing引用目标),学习如何对模型中的特定子图进行精确控制; - 进一步深入:仓库的 docs/source/user_guide/torch_compiler 目录下还有大量专题文档,包括 torch.compiler_dynamo_deepdive.md、torch.compiler_custom_backends.md、torch.compiler_profiling_torch_compile.md 与 torch.compiler_troubleshooting.md,分别覆盖原理深潜、自定义后端、性能剖析与问题排查。
简而言之,torch.compile 的接入成本极低(一行代码),却能借助 TorchDynamo 的图捕获与 TorchInductor 的融合、CUDA graphs 等优化,为从点运算到大型预训练 Transformer 的各类推理负载带来可观的执行效率提升。理解其内核生成与调试手段,是进一步在生产环境落地编译优化、并排查性能瓶颈的起点。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust4.2 K634
MiniCPM5-2BMiniCPM5-2B 是一款面向端侧、本地部署和资源受限场景的 2B 稠密 Transformer,能够达到同尺寸开源模型 SOTA 水平。Markdown300
jforgamejforgame是一个一站式游戏服务器开发框架。包含游戏服务器开发所需要的各种组件,比如网关,socket服务端与客户端,自定义高效消息编解码,游戏热更新,游戏通用工具等等。包含游戏服,跨服,匹配服,后台管理系统等实现,同时提供大量业务案例以供学习。亦可用于其他socket应用,例如及时聊天等。Java101
fizz-gateway-nodeAn Aggregation API Gateway in Java . FizzGate 是一个基于 Java开发的微服务聚合网关,是拥有自主知识产权的应用网关国产化替代方案,能够实现热服务编排聚合、自动授权选择、线上服务脚本编码、在线测试、高性能路由、API审核管理、回调管理等目的,拥有强大的自定义插件系统可以自行扩展,并且提供友好的图形化配置界面,能够快速帮助企业进行API服务治理、减少中间层胶水代码以及降低编码投入、提高 API 服务的稳定性和安全性。Java50
certd开源SSL证书管理工具;全自动证书申请、更新、续期;通配符证书,泛域名证书申请;证书自动化部署到阿里云、腾讯云、主机、群晖、宝塔;https证书,pfx证书,der证书,TLS证书,nginx证书自动续签自动部署JavaScript60
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python280