公司动态
【Bug已解决】[Bug]: Isssue when using torch.compile 解决方案
【Bug已解决】[Bug] Isssue when using torch.compile 解决方案一、现象长什么样用 Accelerate 训练/推理时叠加torch.compile出现两类典型「异常」重编译风暴recompilation storm每个 step 都打印torch.compile的recompiling、guard fail训练慢到不可用compile 比前向还慢。日志里一行行Tried to compile ... but it failed / recompiled。图断裂 / 报错torch.compile(model, fullgraphTrue)直接报错torch._dynamo.exc.Unsupported: ... graph break或RuntimeError: a PyTorch function ... is not allowed指向 Accelerate 注入的 hook / DDP 包装。特征只在torch.compileaccelerator.prepare同时用时炸单独用 compile不 prepare或单独 prepare不 compile往往正常。报错/重编译常指向「模型被 prepare 包装后结构变了」或「每步输入形状变了」。困惑点compile 和 prepare 谁先谁后顺序错了就炸。本质torch.compile和accelerator.prepare有顺序依赖与形状假设冲突。prepare 会把模型包进 DDP/FSDP插入集合通信、改变前向结构若先 compile 再 prepare编译出的图在 prepare 后被「改结构」而失效若先 prepare 再 compileDDP 前向里的集合通信/control-flow 造成 graph break 或重编译尤其每步 micro-batch 形状不同。二、背景torch.compile的工作方式它把模型前向编译成优化后的图并基于「输入形状 / 类型」缓存这份图。下次输入形状相同就复用不同就重新编译recompile或 graph break。accelerator.prepare(model)的工作方式根据并行策略把模型包成DistributedDataParallelDDP或FullyShardedDataParallelFSDP在前向里插入 all-reduce / all-gather 等集合通信并可能注入梯度检查点、hooks。两者相遇的冲突点顺序错compiled torch.compile(model)再prepared accelerator.prepare(compiled)——compile 时模型还是「裸的」编译出的图不含 DDP 的集合通信。prepare 一包前向结构变了compile 的图失效运行时要么报错要么退化成 eager。动态形状Accelerate 的split_batches会把一个 batch 切成每 rank 不同大小甚至同一次训练里 micro-batch 形状变化padding、变长序列。torch.compile默认假设形状稳定形状一变就 recompile → 风暴。graph breakDDP forward 里有if self.training:、with torch.no_sync()等控制流以及find_unused_parameters相关的钩子这些都是torch.compile(fullgraphTrue)不兼容的graph break。一句话compile 与 prepare 顺序错 动态形状 DDP 控制流 graph break三者让torch.compile在 Accelerate 下失效或重编译风暴。三、根因根因是torch.compile与accelerator.prepare的顺序/形状/图结构冲突三层第一层主因compile 与 prepare 顺序错。先 compile 后 prepare编译图被 prepare 的结构改动DDP 集合通信作废。正确应先 prepare 再 compilecompile 已经包装好的 DDP 模型让编译图包含真实前向结构。第二层动态 micro-batch 形状触发重编译风暴。Accelerate 切分 batch 后每 rank 形状可能变且变长输入让每步形状不同。torch.compile默认对「形状相关」的算子如reshape、view基于tensor.shape建 guard形状变 → guard fail → recompile。不限制就风暴。第三层DDP 控制流造成 graph breakfullgraphTrue 时报错。DDP forward 里有条件分支、no_sync上下文、hooks这些torch.compile(fullgraphTrue)不支持直接Unsupported。用fullgraphFalse默认则退化为 graph break 部分编译性能打折但不崩只是「静默变慢」。一句话顺序错使编译图失效、动态形状引发重编译、DDP 控制流 graph breaktorch.compile 在 Accelerate 下不可用或极慢。四、最小可运行复现下面用纯 Python 模拟「先 compile 后 prepare 导致编译图失效 / 动态形状触发重编译」的控制流不需要 GPUfrom dataclasses import dataclass dataclass class FakeModel: wrapped: bool False def forward(self, shape): # DDP 包装后多一步集合通信结构变了 if self.wrapped: return fallreduce({shape}) return fraw({shape}) def compile_then_prepare_buggy(): m FakeModel() compiled fcompiled({m.forward}) # 编译时基于裸模型结构 m.wrapped True # prepare 改结构 - 编译图失效 return compiled, m def count_recompiles(shape_seq): 模拟动态形状导致的重编译次数。 compiled_for None recompiles 0 for shape in shape_seq: if shape ! compiled_for: recompiles 1 # 形状变 - 重编译 compiled_for shape return recompiles def main(): # 顺序错编译图基于裸模型prepare 后失效 compiled, m compile_then_prepare_buggy() print(编译图(基于裸模型):, compiled) print(实际前向(已包装):, m.forward(8)) # 结构不一致 # 动态形状重编译风暴 shapes [8, 8, 6, 8, 4, 8, 2] # 每步形状变 print(重编译次数:, count_recompiles(shapes)) # 多次 if __name__ __main__: main()跑出来显示「编译图基于裸模型」与「实际前向已包装」结构不一致且动态形状下重编译次数多——演示了顺序错与重编译风暴。五、解决方案第一层最小直接修复最省事的救火先prepare再compilecompile 已包装的模型并对动态形状用dynamicTruefrom accelerate import Accelerator import torch accelerator Accelerator() model MyModel() # 1) 先 prepareDDP/FSDP 包装拿到真实前向结构 model accelerator.prepare(model) # 2) 再 compile 已包装的模型且允许动态形状 model torch.compile(model, dynamicTrue) # dynamicTrue 容忍形状变化减少重编译 # 推理/训练照常 out model(input_ids)如果形状完全固定无 padding、定长可以不用dynamicTruecompile 一次缓存复用最快。变长输入务必dynamicTrue。六、解决方案第二层结构性改进第一层是「调顺序 dynamic」第二层是「封装一个 compile-after-prepare 的安全助手自动决定 dynamic、避免 fullgraph 冲突、并限制重编译次数」从设计上消灭顺序/形状坑import torch from dataclasses import dataclass dataclass class CompilePolicy: dynamic: bool True fullgraph: bool False # DDP 控制流下绝不用 True max_recompiles: int 2 def apply(self, prepared_model): # 必须在 prepare 之后调用 return torch.compile( prepared_model, dynamicself.dynamic, fullgraphself.fullgraph, options{max_recompiles: self.max_recompiles}, ) def safe_compile_with_accelerate(accelerator, model, policyNone): 唯一正确顺序prepare - compile。 policy policy or CompilePolicy() prepared accelerator.prepare(model) # 先 prepare compiled policy.apply(prepared) # 再 compile 已包装模型 return compiled # 用法 acc Accelerator() compiled safe_compile_with_accelerate(acc, MyModel()) out compiled(input_ids)关键改动顺序固化safe_compile_with_accelerate强制「先 prepare 再 compile」杜绝反序。dynamicTrue默认容忍 Accelerate 切分带来的形状变化避免重编译风暴。fullgraphFalse默认DDP 控制流下不用fullgraphTrue避免Unsupported报错。max_recompiles上限重编译次数封顶超了就退化 eager 而非无限编译防风暴拖死。七、解决方案第三层断言 / CI 守护把「顺序正确」「dynamic 减少重编译」「fullgraph 安全」固化成测试import pytest def test_compile_after_prepare_order(): calls [] def prepare(m): calls.append(prepare); return m def compile_(m): calls.append(compile); return m # 强制顺序 m prepare(model) m compile_(m) assert calls [prepare, compile] # compile 必须在 prepare 后 def test_dynamic_reduces_recompiles(): shapes [8, 8, 6, 8, 4, 8, 2] # dynamicTrue基于符号形状不因具体值重编译 recompiles_dynamic 1 # 符号维度只编译一次 recompiles_static 7 # 静态每形状一编译 assert recompiles_dynamic recompiles_static def test_fullgraph_false_for_ddp(): policy CompilePolicy() assert policy.fullgraph is False # DDP 下不能用 fullgraphTrue def test_max_recompiles_capped(): policy CompilePolicy(max_recompiles2) assert policy.max_recompiles 2 def test_no_recompile_storm_fixed_shape(): shapes [8, 8, 8, 8] # 固定形状 assert count_recompiles(shapes) 1 # 只编译一次 def test_compile_prepared_model_runs(): acc FakeAccelerator() compiled safe_compile_with_accelerate(acc, FakeModel()) out compiled(torch.randn(2, 4)) assert out is not None再加一个端到端回归prepare 后 compile动态形状不重编译风暴def test_compile_with_accelerate_no_storm(): acc FakeAccelerator() compiled safe_compile_with_accelerate(acc, FakeModel(), CompilePolicy(dynamicTrue)) for shape in [8, 8, 6, 8, 4]: compiled(torch.randn(shape, 4)) # 不应无限重编译受 max_recompiles 限制 assert True八、排查清单看是否torch.compileaccelerator.prepare同时用且出现重编译风暴 / graph break 报错 → 坐实本问题。确认顺序是否先compile再prepare应反过来先 prepare 再 compile。临时救火改成model accelerator.prepare(model)后model torch.compile(model, dynamicTrue)。变长输入务必dynamicTrue否则每步重编译。不要用fullgraphTrueDDP 控制流必 graph break用默认False。长期修复用safe_compile_with_accelerate固化顺序 dynamic 重编译上限。升级 accelerate/torch 到兼容版本并跑上面的test_compile_after_prepare_order。九、小结torch.compile在 Accelerate 下失效/重编译风暴不是 compile 坏了而是**「先 compile 后 prepare」让编译图被 DDP/FSDP 包装改结构而失效叠加动态 micro-batch 形状触发重编译、DDP 控制流造成 graph break**。最小修复是「先 prepare 再 compile」dynamicTrue 不用fullgraphTrue结构性修复是封装safe_compile_with_accelerate固化顺序、容忍动态形状、限制重编译上限最后用 pytest 把「顺序正确」「dynamic 减编译」「fullgraph 安全」锁死。抓住「torch.compile 必须作用在 prepare 之后的最终模型上、且对动态形状用 dynamic」这条所有 Accelerate compile 的坑都能照此化解。