torch.compile 图断裂:先定位重编译,再谈模式选择
从 Dynamo guard、graph break 与动态形状入手,建立可复现的编译诊断和冷启动、稳态验收方法。

1. 先分清「断图」和「反复编译」
接入 torch.compile 后,第一批很慢且随后偶发长尾,团队常立刻切换模式。但模式不能修复 graph break,也不能阻止输入变化触发新 guard。应先回答:一次调用被切成多少段图;同一段图是否因 guard 失效反复编译。
torch.compile 的前端 TorchDynamo 会观察 Python frame,把可追踪的 Tensor 运算提取为 FX graph,再交给 AOTAutograd 和后端(默认通常是 TorchInductor)。遇到无法安全追踪的语义时,Dynamo 可以结束当前图,回到 Python 执行,再从后续可追踪位置开始新图。这是 graph break。它会增加 Python 往返、切断跨边界融合,但若断点位于低频日志或初始化路径,未必值得修。
另一类问题是已有图无法复用。每个缓存条目都有 guard:dtype、device、rank、尺寸、stride、模块属性或 Python 常量都可能成为条件。guard 失效就会编译新变体;p50 正常而 p99 抖动,常见根因是形状或开关制造变体,而非断图数量。
所以诊断顺序应当固定:先用 eager 结果做正确性基线;再看 recompilation 及其 guard failure;然后看 graph break 的位置和频率;最后才比较 compile mode。把所有首次调用都算进平均吞吐,或只测固定 shape 的稳态,都无法回答线上会不会抖。
2. Guard 是编译缓存的契约
下面的函数表面上只有矩阵运算,training 却是 Python 布尔值。若它在调用间变化,Dynamo 通常要为两个分支分别建立缓存条目;batch 尺寸是否被特化,则取决于版本、图和动态形状策略。
import torch
def project(x, weight, training: bool):
y = x @ weight
return torch.nn.functional.dropout(y, 0.1, training=True) if training else y
compiled_project = torch.compile(project)
x = torch.randn(8, 128, device="cuda")
w = torch.randn(128, 256, device="cuda")
compiled_project(x, w, True) # 首次捕获与编译
compiled_project(x, w, True) # 满足 guard 时复用
compiled_project(x, w, False) # Python 常量变化,通常产生另一变体若线上只有 train/eval 两种稳定状态,预热两个变体是合理设计。危险的是把请求 ID、不断变化的长度上限或对象实例传进热点 frame,使变体数随流量增长。优化 guard 的目标不是“零 guard”,而是让变体集合有明确上界。排查时先开日志;PyTorch 2.x 常用 recompiles、guards、graph_breaks 和 dynamic,具体类别应以当前安装版本的 TORCH_LOGS=help 为准。
TORCH_LOGS="recompiles,graph_breaks,dynamic" python serve_bench.py也可在启动早期调用 torch._logging.set_logs(recompiles=True, graph_breaks=True),但参数会随版本演进。重点看失败 guard:是尺寸、模块属性,还是全局开关变化,再决定做 bucketing、动态形状或拆分路径。Dynamo 缓存并非无限,达到重编译上限后可能回退 eager;不同 2.x 版本出现过 cache_size_limit、recompile_limit 等配置名。不要依赖私有配置强行抬高上限;确需调整时应核对当前版本,并同时限制变体数和编译总时间。
3. Graph break 要按热度排序
torch._dynamo.explain 适合在最小复现上汇总捕获结果。现代 PyTorch 2.x 推荐先用函数式调用获得 explain 包装器;旧版本曾支持不同调用形式,因此 CI 中应锁定版本。
import torch
def step(x):
score = x.mean()
if score.item() > 0: # Tensor 值逃到 Python,并驱动控制流
return x.sin()
return x.cos()
report = torch._dynamo.explain(step)(torch.randn(16, device="cuda"))
print("graphs:", report.graph_count)
print("breaks:", report.graph_break_count)
for reason in report.break_reasons:
print(reason)ExplainOutput 字段可能随版本调整,只适合开发诊断。Tensor.item() 常导致标量逃逸;capture_scalar_outputs 也不能把任意 Python if 变成图内分支。两个分支都是安全、纯 Tensor 计算且代价可接受时,可改成 torch.where,但它会先求值两支,再按条件选择结果:原本不会执行的分支也会运行,因此不能机械替换含非法输入、显著大算子或副作用的 Python if。需要惰性执行大块条件计算时再评估 torch.cond;后者在 2.x 中仍有 operand、别名和副作用约束,需核对当前版本并测反向传播。
print、文件 I/O、不支持的扩展、依赖数据的循环以及 list.append 等 Python side effect 都可能断图。让编译区域保持 Tensor 进、Tensor 出,把日志与状态更新留在 eager 外壳。break 应按热度处理:训练 step 内每层发生的优先,每批一次的次之,仅启动一次的通常无需为“单图”重写。
开发阶段可设置 torch.compile(fn, fullgraph=True)。它要求整个调用被捕获成一张图,遇到 graph break 直接报错,因此很适合给“本应纯 Tensor”的 kernel 或 block 写回归测试。fullgraph=True 不是通用加速开关:对包含数据加载、日志和动态宿主逻辑的整套训练循环,它只是把可接受的分段编译变成失败。
4. Dynamic shapes 不是越动态越好
输入长度变化时,静态特化通常能生成更激进的代码,但每个新尺寸可能触发编译;完全动态的图减少变体,却可能丢失特化机会,某些约束还会导致后端无法生成理想 kernel。torch.compile(..., dynamic=None) 在许多 2.x 版本中是默认策略:先静态编译,观察到尺寸变化后尝试自动泛化。dynamic=False 倾向保持静态特化,dynamic=True 则尝试从开始生成更动态的图。这里的“尝试”很重要:并非所有维度和运算都能无条件动态化。
只希望指定维度动态时,可在进入编译函数前用 torch._dynamo.mark_dynamic(x, dim, min=..., max=...);这是版本敏感的底层 API,不能在 trace 内调用,上下界应来自接口契约。工程上通常先做长度 bucketing,把变体限制为可预热的少数桶,再对桶内高基数维度动态化。仍需检查 stride、layout、dtype、device、grad 与 autocast;shape 相同不保证 guard 相同。
5. Python 边界与 custom op
无法被 Dynamo 理解的 C++/CUDA 扩展,若留在 Python 黑盒里,常在其前后形成断图。若该操作确实是稳定热点,应把语义显式注册为 custom operator,让编译系统知道 schema、变异行为和 fake/meta 形状推导,而不是只用装饰器隐藏错误。较新的 PyTorch 2.x 提供 torch.library.custom_op 与 register_fake;它们的签名、autograd 注册方式和编译支持仍有版本差异,下面只展示注册形式,需要按当前版本文档核对。
import torch
@torch.library.custom_op("robot_ops::clamp_norm", mutates_args=())
def clamp_norm(x: torch.Tensor, limit: float) -> torch.Tensor:
# 示例实现;真实项目可在这里调用已加载的 C++/CUDA 内核
norm = x.norm(dim=-1, keepdim=True).clamp_min(1e-6)
return x * (limit / norm).clamp(max=1.0)
@clamp_norm.register_fake
def _(x, limit):
return torch.empty_like(x)
def block(x):
return torch.relu(clamp_norm(x, 2.0))
compiled_block = torch.compile(block, fullgraph=True)
# 对每种受支持的 device、dtype 和关键 shape 准备代表性输入
torch.library.opcheck(
clamp_norm,
(torch.randn(4, 128, device="cuda"), 2.0),
)示例中的 clamp_norm 函数体本身全是普通 PyTorch Tensor 运算,实际代码应直接让 Dynamo 追踪,而不应为了“单图”把这类函数包装成 custom op。custom op 对 Dynamo 和后端都是 opaque 的:它能保留在捕获图中、消除外围 Python graph break,却仍是算子级融合边界,Inductor 不会进入函数体把内部运算与相邻算子融合。这里用 Tensor 实现只是让注册契约可读;真实用例通常是在函数体中调用无法追踪的外部 C++/CUDA kernel。
register_fake 只描述输出元数据,不证明 kernel 正确。torch.library.opcheck 会检查 schema、变异声明、FakeTensor 元数据和 AOT dispatch 等注册契约,但同样不替代数值测试;应对所有支持的 device、dtype、边界 shape 与非连续 layout 提供样例。涉及梯度还要按当前 API 注册 autograd,并另做 gradcheck、eager/compiled 对齐。不要把 torch.compiler.allow_in_graph 当普通修复,它绕过部分安全检查,可能把显式断图变成静默错误。低频扩展保留 eager 边界往往更经济;只有消除断图、支持 fullgraph/导出或把外部 kernel 纳入编译调度的收益覆盖注册维护成本时,custom op 才有回报。
6. Regional compilation 控制冷启动半径
大型模型直接 torch.compile(model) 会让首次捕获覆盖很大区域,冷启动时间、峰值编译内存和失败定位范围都随之扩大。若 Transformer 的同构 block 重复数十次,可以只编译 block 或一组 block,让调用方和数据相关路由保持 eager。这通常被称为 regional compilation;它是一种边界设计,不是要求某个固定的专用 API。
不要假设各 block 必然共享代码缓存,应从日志确认。区域过小会增加 Python 调度并失去跨 block 融合,过大又恢复漫长冷启动;边界宜落在重复、计算密集且契约稳定的模块。
这时再比较模式:reduce-overhead 常面向小 batch 与 launch 开销,可能利用 CUDA Graphs,也可能增加为重放保留 workspace 的内存;max-autotune 付出更多编译成本搜索矩阵乘、卷积等实现,并且在部分 PyTorch 2.x GPU 配置中默认连带启用 CUDA Graphs。若要把 autotune 收益与 graph replay 的收益、内存代价分开,应在当前版本支持时加入 max-autotune-no-cudagraphs。默认模式仍作为基线。模式集合和具体开关会随版本、后端与设备变化,torch._inductor.list_mode_options() 可用于当前环境检查,但 _inductor 是内部命名空间,不能把某一版本的展开选项当成长期契约。所有模式必须共享同一 shape 和预热协议。
7. 把冷启动与稳态分开计量
一个可复现 benchmark 至少包含三段:新进程中的首次调用;覆盖允许 shape/状态集合的预热;预热后的长时间稳态。首次调用记录端到端时延、编译日志和进程峰值内存;预热记录产生了多少图变体以及总耗时;稳态记录吞吐、p50/p95/p99、GPU 利用率,并持续观察是否还有 recompile。
GPU 微基准要在计时边界调用 torch.cuda.synchronize(),并把输入生成移出计时;训练还应分测 forward、backward 和 optimizer step。冷启动超预算时可在接流量前预热有限 shape 桶;请求分布无有限上界时则必须测试动态策略与缓存压力。多 worker 还要确认每个进程是否独立编译,以及滚动发布的并发冷启动内存。
8. 验收结论应能解释每一次编译
上线前保留 shape、dtype、layout、grad mode 和关键 Python 状态的输入矩阵,每一格都跑 eager/compiled 对齐;训练还要比较梯度和若干步参数更新。融合会改变舍入顺序,容差应按 dtype 与任务设定。
性能侧的通过条件应写成可观测契约:
- 冷启动时间、峰值内存与变体数在预算内;
- 预热覆盖 shape 桶和 train/eval、autocast 状态,之后没有未解释的 recompile;
- 热路径 break 有取舍记录,稳态按真实请求报告延迟分位数和吞吐;
- custom op 通过
opcheck、eager/compiled 数值对齐与必要的梯度测试; - 固定 PyTorch、CUDA、Triton、GPU 和编译配置。
最终要优化的不是 explain 报告里的“零断图”,而是一个稳定、可解释的执行系统:有限且可预热的编译变体,热路径上足够大的融合区域,冷启动符合发布预算,稳态不再被 recompile 长尾打断。先用 guard failure 解释重编译,再按热度处理 graph break,最后在相同边界上选择 mode,才能把 torch.compile 从一次性的跑分开关变成可验收的生产能力。
相关
也可以看看
- ·3 分钟阅读
PyTorch SDPA 的两个静默坑:布尔 mask 与 eval dropout
拆清 scaled_dot_product_attention 的 keep-mask 语义、广播维度和 dropout_p 行为,用最小张量测试阻止注意力泄漏。
- ·6 分钟阅读
CE target long dtype 断言
PyTorch 广播让错误 shape 的 loss 能跑起来——silent bug 产生错误梯度;写 forward 时 assert shape 或用 torchdim 习惯。
- ·16 分钟阅读
KV cache 分页与前缀复用:吞吐上去后正确性怎么验
沿 block table、前缀哈希与抢占回收拆开缓存一致性风险,说明错页复用为何表现为「偶发胡话」;用确定性前缀集、强制抢占和逐 token 对照验收。
johan's blog