ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch 实践指南:在何处应用 torch.compile——编译粒度、应用位置与最佳实践

PyTorch 实践指南:在何处应用 torch.compile——编译粒度、应用位置与最佳实践 PyTorch 实践指南在何处应用 torch.compile——编译粒度、应用位置与最佳实践【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchtorch.compile是 PyTorch 的即时编译技术能够在不改变模型代码语义的前提下通过图捕获与内核生成显著提升训练和推理性能。但编译哪里、以什么粒度编译直接决定了编译收益、首次编译耗时与代码可维护性。本文以 PyTorch 官方用户指南 为骨架结合当前仓库源码系统讲解torch.compile的推荐应用位置、nn.Module.compile()与函数式torch.compile(model)的区别以及区域编译/层级编译在降低编译耗时中的应用。读完本文你将能够为推理、训练、分布式DDP/FSDP场景选择正确的编译粒度并写出可落地的编译代码。核心原则编译最高且不引发问题的函数官方指南给出的总原则只有一条将torch.compile应用到不会引发过多问题的最顶层函数。这里的问题包括图捕获失败导致的 graph break、与分布式包装器的交互异常、不必要的重编译等。典型的高层应用位置有三种训练train或推理eval步即包含 optimizer 操作但不包含外层循环的那个函数顶层nn.Module直接对模型实例调用.compile()某些子nn.Module当模型过大、结构重复或首次编译时间过长时退而求其次编译其中的子模块。之所以推荐尽可能高层是因为torch.compile的收益来自更大范围的图融合fusing与内核生成编译区域越大跨算子优化的机会越多Python 层调度开销被消除得也越彻底。而不要过高的约束则来自两个现实一是把整个训练循环含数据加载、日志、Python 控制流都塞进编译区会因大量 graph break 而让收益大打折扣二是分布式包装器DDP/FSDP与编译器的交互并不完美见下文专门讨论。三个标准应用场景推理、训练、DDP推理场景直接编译模型推理是torch.compile最简单的应用形态——只需在进入推理循环之前对模型调用.compile()即可# inference model ... model.compile() for _ in range(N_ITERS): inp ... out model(inp)注意编译动作model.compile()必须放在循环之外torch.compile是 JIT 编译器首次调用编译后的模型时会付出一次编译成本此后命中的是编译缓存compile cache只有把编译放在循环外才能让N_ITERS次迭代全部复用同一份编译产物。训练场景编译包含 optimizer 的 train 步训练场景推荐把编译范围覆盖到含 optimizer、不含数据加载与循环的 train 函数上# training model ... opt torch.optim.Adam(model.parameters()) torch.compile def train(mod, data): opt.zero_grad(True) pred mod(data[0]) loss torch.nn.CrossEntropyLoss()(pred, data[1]) loss.backward() opt.step() for _ in range(N_ITERS): inp ... train(model, inp)这个示例展示了三个值得注意的细节编译的是函数而非模块用torch.compile装饰train(mod, data)模型以参数形式传入。这正是下节要展开的函数式编译合法形态之一优化器步进被纳入编译区opt.zero_grad(True)、loss.backward()、opt.step()都在编译边界之内。Dynamo 能够跟踪 PyTorch 优化器的这些操作使梯度清零、反向传播与参数更新共享同一份编译后的执行计划循环留在编译区外外层for _ in range(N_ITERS)保持为普通 Python 循环每个 epoch/step 调用一次编译后的train函数。分布式场景编译传给包装器的内层模块指南特别指出torch.compile不能很好地处理 DDP、FSDP 这类分布式包装模块torch.compilespecifically doesn’t handle distributed wrapper modules like DDP or FSDP very well。因此推荐的写法是编译传给包装器的内层模块而不是包装器本身# DistributedDataParallel model ... model.compile() model_ddp DistributedDataParallel(model, ...) for _ in range(N_ITERS): inp ... out model_ddp(inp)从当前仓库源码可以印证这一建议的合理性在 torch/distributed/fsdp/_fully_shard/_fsdp_common.py 中FSDP 定义了一个_dynamo_disable辅助函数其实现直接包装torch._dynamo.disable(...)并把它应用到_fsdp_param_group.py、_fsdp_state.py、_fully_shard.py中的多处关键方法上见 torch/distributed/fsdp/_fully_shard/_fsdp_state.py 中的多个_dynamo_disable装饰。这说明 FSDP 自身的状态机、参数分组逻辑在编译器前端Dynamo看来是不透明的若直接编译 FSDP 包装后的模块这些被 disable 的边界会引入 graph break 或导致追踪结果与预期不符。把编译粒度下移到内层model让分布式通信发生在编译区域之外是最稳妥的做法。torch.compile(model)与model.compile()为什么推荐方法形式指南用专门一节澄清了一个极易踩坑的差异由于torch.compile与nn.Module实例交互时存在一些微妙行为如果要把某个模块作为顶层函数编译建议使用nn.Module实例自带的.compile()方法# DO NOT DO THIS model MyModel() model torch.compile(model) model(inp) # DO THIS model MyModel() model.compile() model(inp) # this is also acceptable torch.compile def fn(model, inp): return model(inp) model MyModel() fn(model, inp)为什么DO NOT DO THIS的写法不推荐关键在nn.Module.compile()的实现细节。查看 torch/nn/modules/module.pydef compile(self, *args, **kwargs) - None: Compile this Modules forward using :func:torch.compile. This Modules __call__ method is compiled and all arguments are passed as-is to :func:torch.compile. See :func:torch.compile for details on the arguments for this function. self._compiled_call_impl torch.compile(self._call_impl, *args, **kwargs)也就是说model.compile()实际编译的是Module._call_impl方法__call__的实现主体并把编译结果缓存在实例属性_compiled_call_impl上。这意味着编译目标是方法而不是函数对象_call_impl是绑定在实例上的方法Dynamo 在追踪nn.Module时对_call_impl有专门的处理逻辑。在 torch/_dynamo/variables/nn_module.py 中可以找到对应证据Dynamo 会检查mod._call_impl是否仍是未被包装的原始torch.nn.Module._call_implunpatched_nn_module_call_impl见 torch/_dynamo/utils.py若是则可以直接把_call_impl短路short-circuit到forward从而内联子模块调用嵌套模块调用会被正确追踪指南明确说嵌套模块调用会被正确追踪此时无需再对每个子模块调用.compile()Nested module calls will be traced correctly - there is no need to call.compile()in that case。这与上述内联机制一致——编译顶层模块时其内部所有子模块的forward都会在符号追踪中展开子模块无需各自编译。而model torch.compile(model)之所以被标记为 DO NOT DO THIS是因为函数式torch.compile作用于模块实例时模块实例的包装/替换语义与nn.Module内部对_call_impl、forward的调度逻辑包括 hook、_compiled_call_impl缓存等交互存在边界情况容易导致编译行为不符合直觉例如 graph 捕获粒度不一致、或某些属性访问被意外包装。如果你偏好函数式写法官方给出的可接受替代是torch.compile装饰一个接收model的普通函数如示例中的fn(model, inp)——既保持了函数式风格又避开了直接包装模块实例的问题。另外需要说明torch.compiler.compile与torch.compile是同一入口torch/compiler/init.py 中的compile(*args, **kwargs)只是把参数原样转发给torch.compile因此上文的编译位置原则对两种写法均适用。编译子模块用区域编译与层级编译降低首次编译耗时如果整个模型太大一次全量编译的冷启动耗时可能从数秒到数分钟超大模型甚至可达数十分钟。指南给出了另一个实用建议对更小的重复区域例如单个 transformer block而不是整个模型应用torch.compile可以大幅缩短编译时间并引导读者查看 Reducing Compile Time 文档中的 regional and hierarchical compilation 小节。结合 programming_model.reducing_compile_time.md 的内容两种降低编译耗时的细分策略是区域编译Regional compilation单独编译模型中反复出现的较小区域如单个 transformer block让该编译区域在每次重复出现时复用同一份编译产物。由于区域只编译一次而不是每个重复处各编译一次对于由大量相同 block 堆叠而成的模型典型如 LLM冷启动编译时间可显著下降层级编译Hierarchical compilationtorch.compiler.nested_compile_region在一次torch.compile调用内部用torch.compiler.nested_compile_region标记重复的结构单元同样是 transformer layer 这类单元。编译器首次遇到该区域时生成优化代码之后每次出现时直接盖章复用stamp out已编译代码而非重新编译。与区域编译不同层级编译不需要重构你应用torch.compile的方式它天然工作在单个torch.compile之内同时它是安全的——若新输入的条件shape、dtype、device、stride、globals 等会使缓存的区域失效编译器会透明地重编译以保持正确性在torch.compile上下文之外该装饰器则是 no-op。从torch.compiler的公开 API 列表torch/compiler/init.py可以看到nested_compile_region与compile、config、reset、disable、allow_in_graph等并列导出是编译编程模型中一等公民的组成能力。选择编译粒度的决策清单综合指南与源码实践中可按以下顺序决策先考虑编译整个训练/推理步函数含 optimizer、不含循环这是收益最大的默认选项涉及 DDP/FSDP 时将编译下移到传给包装器的内层模型让分布式通信留在编译区外需要编译顶层nn.Module时优先model.compile()编译_call_impl避免model torch.compile(model)的写法偏好函数式则用torch.compile装饰接收 model 参数的函数冷启动编译时间成为瓶颈时先测量torch._dynamo.utils.compile_times()、TORCH_COMPILE_DEBUG1、tlparse/TORCH_TRACE见 programming_model.observability再用区域编译或nested_compile_region层级编译把编译粒度细化到重复结构单元编译完就固定输入形态避免因 shape 变化触发不必要的重编译重编译的排查见 Dealing with Recompilations。遵循这一决策链你就能在编译收益最大化与编译成本、分布式兼容性、代码可维护性之间取得平衡让torch.compile在推理、训练与分布式场景中都稳定落地。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表