torch.compile能正常跑完,不代表模型已经形成稳定的编译图。默认fullgraph=False时,遇到无法追踪的 Python 代码,Dynamo 会结束当前图,执行那段代码,再从后面继续追踪。程序结果仍然正确,但图被切碎后,算子融合机会减少,编译开销还可能反复出现。
先准备一个最小复现环境。以下步骤不依赖特定模型,CPU 也能观察图断裂:
import torch @torch.compile def step(x): y = torch.sin(x) + torch.cos(x) if y.sum().item() > 0: y = y * 2 return torch.relu(y) x = torch.randn(1024) for _ in range(5): step(x)这里的.item()把 Tensor 值取回 Python,后面的分支依赖运行时数据。PyTorch 官方排障文档把数据依赖控制流列为常见 graph break 来源。默认模式下代码通常不会直接报错,所以只看输出和总耗时,很难知道编译区域已经断开。
第一步开启细粒度日志:
$env:TORCH_LOGS="graph_breaks,recompiles" python .\demo.pyLinux 或 macOS 可以写成:
TORCH_LOGS="graph_breaks,recompiles" python demo.pygraph_breaks会给出用户代码位置和原因,recompiles记录 guard 失败导致的重新编译。先解决断裂,再处理重编译,不要把两者混在一起。断裂是一次执行被切成多个图;重编译是已有编译结果不再适合新的输入或状态。
第二步用fullgraph=True把静默断裂变成错误:
@torch.compile(fullgraph=True) def step(x): ...这种模式要求整个函数可捕获,遇到第一个 graph break 就停止。它适合缩小复现范围,不一定适合直接作为生产配置。大型模型可能包含日志、预处理、第三方算子和难以追踪的辅助代码,强求一个图会让迁移成本很高。
针对上面的数据依赖分支,可以把条件改写成 Tensor 运算,例如使用torch.where,前提是两个分支都适合计算,而且不会引入不可接受的额外开销。不要为了消除日志中的每条断裂,改变原有数值语义。官方文档也提醒,不同 graph break 的代价不同。发生在forward中间的断裂通常比发生在预处理边缘更值得处理。
第三步检查动态形状。批大小、序列长度或图像尺寸变化时,guard 可能失败并触发重新编译。TORCH_LOGS="guards,recompiles,dynamic"能看到保护条件和动态形状处理。先用固定输入尺寸跑基线,再逐步放开一种维度。一次改变多个维度,日志会很快失去可读性。
如果某段代码频繁断裂且没有编译收益,可以明确跳过:
@torch.compiler.disable def preprocess(x): print(x.shape) return x这比让编译器每次尝试、失败、恢复更可控。recursive=False可允许被禁用函数内部的其他调用继续追踪,但使用前要确认调用边界。torch._dynamo.config.suppress_errors = True会在编译错误后回退 eager,官方排障页不建议把它当长期解决方案,因为它容易把真实问题藏起来。
第四步区分首次编译时间和稳态推理时间。至少预热多轮,单独记录第一次调用、后续稳定调用和输入形状变化后的调用。GPU 测量要在计时点同步 CUDA,否则 Python 计时可能只测到异步提交。比较 eager 和 compile 时使用同一批输入、相同精度设置,并确认输出误差在允许范围内。
第五步使用 profiler 看编译区域。官方 profiling 文档说明,时间线中的Torch-Compiled Region可以帮助判断编译是否真正覆盖关键路径。若大量小区域被 Python 间隙隔开,即使没有异常,融合空间也会有限。先修复最热路径中重复出现的断裂,不要追求日志绝对为零。
上线前再做两组回归。一组覆盖真实尺寸分布,观察编译缓存和显存是否稳定;另一组覆盖异常输入,确认回退路径不会突然放大延迟。torch.compile是运行时优化,不是给函数加一个装饰器就结束。能跑只是起点,图断裂位置、重编译次数和稳态时间才说明它有没有工作。
复现脚本、环境版本和日志开关应一起保存。编译器行为会随 PyTorch、CUDA 与后端变化,没有版本信息的性能结论很难复查,也不适合直接用于下一次升级。