torch.compile 与 Inductor 编译器
torch.compile 是 PyTorch 2.x 的核心特性,通过 JIT 编译将动态图转换为优化的静态图,显著提升推理和训练速度。
使用方式
python
import torch
model = MyModel().cuda()
# 一行代码启用编译
compiled_model = torch.compile(model)
# 训练时自动优化
for batch in dataloader:
loss = compiled_model(batch)
loss.backward()
# 推理加速
output = compiled_model(input)Inductor 后端
Inductor 是 torch.compile 的默认后端,生成 Triton Kernel:
- Dynamo 捕获:将 Python 代码转换为 FX Graph
- 图优化:融合算子、消除冗余
- 代码生成:生成 Triton Kernel 或 C++ 代码
编译开销
首次编译会花费较长时间(秒级),后续调用使用缓存的编译结果。建议在训练前预热模型。