torch.compile
进阶PyTorch 2.0 起提供的一行式模型编译加速接口。
torch.compile 是 PyTorch 2.0(2023 年)引入的编译功能。平时 PyTorch 是「逐个算子立即执行」,每一步都有 Python 调度开销,也无法跨算子优化。torch.compile 用 TorchDynamo 在运行时抓取 Python 代码里的计算图,再交给默认后端 TorchInductor 做算子融合、生成 Triton/C++ 内核(内核即在 GPU/CPU 上实际执行的底层函数)。用法通常只需把模型包一层,首次调用会花时间编译,之后训练和推理变快。在具身领域它常用来压低 VLA 等大模型的单步推理延迟,也可配合 CUDA Graph 减少内核启动开销。
例子部署扩散策略时写 policy = torch.compile(policy, mode=「max-autotune」),预热几次后每步去噪的耗时下降。
- 也叫
- PyTorch 2 编译
- 相关
- PyTorch、推理延迟、CUDA Graph、算子、TensorRT、推理部署
- 来源
- torch.compile — PyTorch documentation
PyTorch 2.x overview