torch.compile:一行代码的编译加速
本教程共 60 篇 · 第 46 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:理解 PyTorch 的编译加速思路,学会用
torch.compile一键加速模型,并搞懂图中断和编译缓存的含义。
为什么要编译
先回顾 eager 模式(即时执行模式)的流程:PyTorch 默认一行一行执行代码,每行张量运算立刻算出结果。好处是灵活、好调试,坏处是慢。慢在哪?每个算子都要经过 Python 解释器,还有大量中间结果在显存里进进出出。
打个比方。eager 模式像逐句翻译:听一句翻一句,翻完就忘。编译模式(compilation)像先听完整场演讲再动笔:整体把握上下文,能合并的句子合并,能省的动作省掉。
开销具体在哪?我数给你看:每个算子调用要经过 Python 解释器分发,几百个算子就是几百次;每次运算还要分配中间张量,写进显存再读出来;GPU 上每个内核启动也有固定成本。小算子一多,这些开销加起来比计算本身还贵。torch.compile 把这一堆小算子揉成几个大内核,开销自然就压下去了。
torch.compile 干的就是这件事:把你的 Python 代码抓成一张计算图,然后整体优化、编译成高效内核。而且——只需要改动一行代码。
一行代码的用法
torch.compile 是个装饰器,也能当普通函数调用:
import torch
def foo(x, y):
a = torch.sin(x)
b = torch.cos(y)
return a + b
opt_foo = torch.compile(foo) # 方式一:包装函数
print(opt_foo(torch.randn(3, 3), torch.randn(3, 3)))
或者用装饰器语法:
@torch.compile
def foo(x, y):
a = torch.sin(x)
b = torch.cos(y)
return a + b
训练时最常用的场景是编译整个模型,两种写法等价:
model = MyModel()
model.compile() # 方式一:调用 .compile() 方法
model = torch.compile(model) # 方式二:直接包装模块
编译是递归的:顶层函数里调用的子函数,也会一起被编译。编译后的模块还是个 nn.Module,能正常参与训练、保存和加载,其他代码不用改。官方推荐直接调用 model.compile(),语义更清楚。
torch.compile 还有 mode 参数,比如 mode="reduce-overhead" 用 CUDA 图进一步压 Python 开销,mode="max-autotune" 编译时自动选最快实现,这个下一章细讲。
第一次调用慢,别慌
编译需要时间。第一次调用编译后的函数,往往比 eager 还慢很多,这是正常的。编译结果会被缓存,之后调用就快了。
拿官方教程的测速看:一个 4096×4096 的矩阵操作,第一次跑 569ms(编译开销),后续稳定在 0.36ms。eager 模式中位数 0.87ms,编译后快了约 2.4 倍。加速主要来自两点:减少 Python 开销、减少显存读写。
真实模型上效果更直观。官方教程用 DenseNet-121 在 GPU 上测:推理从 17.6ms 降到 8.4ms,约 2.1 倍;训练从 51.2ms 降到 20.6ms,约 2.5 倍。注意第一次训练迭代花了 166 秒,几乎全在编译。更多模型的加速对比,可以看官方的 TorchInductor 性能仪表板。
Note加速效果和模型结构强相关。计算密集的大模型收益明显;小算子、Python 逻辑多的代码,图中断多,收益可能很小。
图中断:编译的减速带
torch.compile 要抓计算图,但 Python 太灵活,不是所有代码都能抓。遇到搞不定的代码,它会把图打断:能编译的部分编译,剩下的退回 Python 执行。这个现象叫图中断(graph break)。
最常见的触发点是依赖数据的控制流,比如 if x.sum() > 0 这种条件取决于张量值的分支。图中断不是错误,只是优化机会变少,代码仍能正常运行。
运行时它会先跑编译好的图,遇到断点交回 Python 执行,再进下一个图。图被切得越碎,Python 开销回得越多。容易触发图中断的代码,除了数据依赖的分支,还有动态的 Python 容器操作、外部库调用。经验法则:能写成张量运算的,就别写成 Python 循环。
想看到底断在哪,开日志:
torch._logging.set_logs(graph_breaks=True)
想强制不允许中断,用 fullgraph=True,一旦出现图中断直接报错:
opt = torch.compile(foo, fullgraph=True)
报错后可以改代码:数据依赖的分支用 torch.cond 表达,torch.compile 就能把它们也纳入图中。类似地,torch.while_loop 可以处理循环。
和 TorchScript 的关系
老读者可能听过 TorchScript。它也能把模型转成图,但体验差不少:torch.jit.trace 遇到数据依赖的控制流会静默出错,torch.jit.script 要求改代码、写类型注解。
torch.compile 两者都不用,原样 Python 代码直接编译。与 TorchScript 相比,它更灵活,错误信息也更友好。当年我在 TorchScript 上被类型注解折腾过好几回,换到 torch.compile 之后,同样的代码一行没改就编译通过了。
WarningTorchScript 已弃用(
jit.trace和jit.script不再活跃开发)。官方推荐:要加速,用torch.compile;要导出部署,用torch.export。第 50 章会讲导出。
编译缓存:第二次就快了
torch.compile 的缓存是自动的。TorchDynamo、Inductor、Triton 各层都有磁盘缓存,编译过的图、生成的内核都会存下来,下次直接复用。默认目录在 TORCHINDUCTOR_CACHE_DIR 指定的位置,Linux 下类似 /tmp/torchinductor_用户名。
进程间也能共享缓存。PyTorch 2.4+ 提供端到端缓存(Mega-Cache)API:
# 机器 A:编译后保存
result = opt_fn(a, b)
artifact_bytes, cache_info = torch.compiler.save_cache_artifacts()
# 机器 B:加载后直接命中缓存
torch.compiler.load_cache_artifacts(artifact_bytes)
缓存会校验 PyTorch 和 Triton 版本,GPU 还会校验设备型号,版本不一致不会复用,别担心脏缓存。
另外注意:缓存是按输入形状存的。换个 batch size 或输入长度,属于新形状,会触发重新编译,第一次又慢一下。训练时 batch size 固定,基本没影响;推理服务输入长度多变的话,编译次数会明显变多,这是正常的。
Tip改了模型代码但总觉得「没生效」?用
torch._dynamo.reset()清掉内存缓存,必要时删掉TORCHINDUCTOR_CACHE_DIR下的目录。缓存占点磁盘空间也别慌,删了无非是下次重新编译。
小结
torch.compile 的价值在于:几乎不改代码,白拿加速。记住三个要点:第一次调用慢是编译开销;图中断减少优化机会但不报错;缓存让第二次运行快起来。编译前后数值基本一致,个别场景有微小舍入差异,不影响使用。下一章深入 Inductor,看看编译后代码长什么样、怎么调优。