ARTICLE DETAIL

资讯详情

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

PyTorch Compiled Autograd 实战指南:用 torch.compile 捕获更大的反向传播图

PyTorch Compiled Autograd 实战指南:用 torch.compile 捕获更大的反向传播图 示例工程【免费下载链接】tutorialsPyTorch tutorials.项目地址https://gitcode.com/gh_mirrors/tuto/tutorials点击查看免费下载导读本文基于 compiled_autograd_tutorial.rst 展开系统讲解 PyTorch 2.4 引入的 Compiled Autograd 机制——它通过直接接入 autograd 引擎在运行时捕获比 AOTAutograd 更完整、更大的反向传播图从而消除前向图断点向反向传播的传导、并纳入 backward hooks。读完本文你将掌握 Compiled Autograd 的启用方式全局开关与上下文管理器、日志观测手段TORCH_LOGS、与torch.compile组合使用时的图结构行为以及两类最常见的重编译触发原因与排查方法。背景为什么需要捕获更大的 backward graphtorch.compile的完整编译流水线包含三个核心组件TorchDynamo前端字节码捕获、AOTAutograd前向/反向图切分与预编译、Inductor后端代码生成。其中AOTAutograd 负责**提前ahead-of-time捕获 backward graph但这种捕获是部分partial**的存在两类结构性限制前向图断点会传导到反向只要 forward 中出现 graph breakbackward 也会随之被切成多段损失优化机会反向钩子backward hooks不被捕获tensor.register_hook等钩子逻辑停留在 Python 层无法进入编译图导致钩子路径无法被内核融合优化。Compiled Autograd编译自动微分正是针对这两个痛点设计的torch.compile扩展。它不是替代 AOTAutograd而是在其之上增加一层运行时直接与 autograd 引擎集成捕获完整的 backward graph。凡是模型具备上述两种特征forward 中有 graph break、依赖 backward hooks都值得尝试开启 Compiled Autograd并可能观察到更好的性能。从源码目录看本仓库的 index.rst 与 compilers_index.rst 均将该教程归入 Model-Optimization 专题与 torch_compile_tutorial.pytorch.compile入门、torch_export_tutorial.pytorch.export导出等文档构成完整的编译器工具链学习路径。Compiled Autograd 自身的代价捕获更大图的同时Compiled Autograd 也引入新的开销与风险backward 起始阶段新增运行时开销每次loss.backward()进入编译引擎时需要先做缓存查找cache lookup更容易触发重编译与 graph break由于捕获范围更大Dynamo 在后向图上遇到的重编译和断点概率相应上升。需要特别说明的是Compiled Autograd 仍处于活跃开发阶段尚未与所有现有 PyTorch 特性完全兼容针对特定特性的最新支持状态请以官方 Compiled Autograd 状态页为准原文档给出了对应的 Google Docs 链接此处不展开。此外本文所有示例的默认后端以仓库教程为准部分示例显式指定backendaot_eager或backendeager以便于观察图结构。环境准备与前提PyTorch 2.4Compiled Autograd 在 2.4 中引入建议先完成仓库中的 torch_compile 入门教程理解torch.compile、graph break、fullgraphTrue等基础概念建议通读 PyTorch 2.x 中 TorchDynamo 与 AOTAutograd 两个组件的设计思路。本文示例全部基于一个极简模型输入 10 维向量经单个线性层映射为 10 维输出。它足够小能让图结构一目了然同时不失代表性。import torch class Model(torch.nn.Module): def __init__(self): super().__init__() self.linear torch.nn.Linear(10, 10) def forward(self, x): return self.linear(x)基本用法两行配置开启 Compiled Autograd启用 Compiled Autograd 的推荐方式是在调用torch.compile之前将全局配置项置为Truemodel Model() x torch.randn(10) torch._dynamo.config.compiled_autograd True torch.compile def train(model, x): loss model(x).sum() loss.backward() train(model, x)代码流程拆解如下创建Model实例并用torch.randn(10)生成随机的 10 维输入张量x定义训练函数train内部依次计算loss model(x).sum()并调用loss.backward()整个函数被torch.compile装饰调用train(model, x)触发编译执行。调用链逐层拆解从 Python 解释器到 Compiled Autograd 引擎当train(model, x)被执行时整个调用链依次经过以下环节Python 解释器 → Dynamo由于函数带torch.compile装饰Python 解释器将调用转发给 TorchDynamoDynamo 字节码捕获Dynamo 拦截 Python 字节码模拟其执行并记录操作形成一张计算图AOTDispatcher 切分前向/反向AOTDispatcher 会禁用 hooks调用 autograd 引擎计算model.linear.weight与model.linear.bias的梯度并记录到图中随后通过torch.autograd.Function将train的前向与反向实现重写为可编译形式Inductor 代码生成Inductor 为 AOTDispatcher 的前向/反向生成一份优化实现对应的函数Dynamo 注册优化函数Dynamo 将优化后的函数设置为 Python 解释器接下来要执行的代码解释器执行前向执行loss model(x).sum()解释器执行 backward执行loss.backward()进入 autograd 引擎由于第 1 步设置了torch._dynamo.config.compiled_autograd Trueautograd 引擎将请求路由到Compiled Autograd 引擎Compiled Autograd 捕获全图Compiled Autograd 计算weight、bias的梯度并记录到一张图中包括沿途遇到的所有 hooks在此过程中它会记录 AOTDispatcher 之前重写过的 backward随后生成一个对应loss.backward()完整追踪实现的新函数并在inference 模式下用torch.compile执行它递归编译同样的步骤递归作用于 Compiled Autograd 图但此时 AOTDispatcher 不再需要对图进行切分。从这张流程图可以清晰看出AOTAutograd 的职责切分、预编译部分反向被完整保留Compiled Autograd 只是在运行时把整段 backward 补全成一个更大的图。这解释了为何该机制能同时解决“前向断点传导”与“hooks 不被捕获”两个问题。用 TORCH_LOGS 观察 Compiled Autograd 的图复现上面示例后可以通过环境变量让 Compiled Autograd 把图打印到stderr。该机制基于 PyTorch 的TORCH_LOGS日志体系torch._logging仓库教程 torch_compile_tutorial.py 中即演示了torch._logging.set_logs(graph_codeTrue)的用法# 仅打印 compiled autograd 图 TORCH_LOGScompiled_autograd python example.py # 打印带有更多张量元数据与重编译原因的图以性能为代价 TORCH_LOGScompiled_autograd_verbose python example.py观察日志时注意一个命名规律图中部分节点的名字带有aot0_前缀它们对应 AOTAutograd 中预先编译的 backward graph 0 中的节点。例如aot0_view_2表示 AOT backward graphid0中的view_2节点——这正是“Compiled Autograd 复用 AOTAutograd 已捕获部分”的直接证据。下图展示了compiled_autograd_verbose日志的完整输出。红色框内封装的正是没有 Compiled Autograd 时torch.compile能捕获到的 AOT backward 图其余部分则是 Compiled Autograd 额外补全的内容重要提示上图展示的是将在其上调用torch.compile的输入图而非优化后的图。Compiled Autograd 本质上生成一段“未经优化的 Python 代码”用来表示整个 C autograd 执行过程真正的优化内核融合、算子选择仍由后续的torch.compile完成。前向与反向使用不同编译配置由于前向由 AOTAutograd 处理和反向由 Compiled Autograd 处理是两次独立编译你完全可以用不同的编译器配置分别对待它们。例如即使前向存在 graph break也可以要求反向必须是完整单图fullgraph。方式一函数内分别编译def train(model, x): model torch.compile(model) loss model(x).sum() torch._dynamo.config.compiled_autograd True torch.compile(lambda: loss.backward(), fullgraphTrue)()这里前向用默认配置编译model反向则通过torch.compile(..., fullgraphTrue)强制整段loss.backward()编译成单图任何 graph break 都会直接抛错。方式二上下文管理器更推荐的写法是上下文管理器torch._dynamo.compiled_autograd.enable(...)它会对作用域内的所有 autograd 调用生效def train(model, x): model torch.compile(model) loss model(x).sum() with torch._dynamo.compiled_autograd.enable(torch.compile(fullgraphTrue)): loss.backward()两种方式等价但上下文管理器将启用范围显式限定在代码块内避免了全局开关对其他代码路径的隐式影响也更适合嵌入大型训练脚本。Compiled Autograd 如何解决 AOTAutograd 的两大限制场景一前向图断点不再必然导致反向图断点先看一个在 forward 中人为制造两次 graph break 的例子torch.compile(backendaot_eager) def fn(x): # 1st graph temp x 10 torch._dynamo.graph_break() # 2nd graph temp temp 10 torch._dynamo.graph_break() # 3rd graph return temp.sum() x torch.randn(10, 10, requires_gradTrue) torch._dynamo.utils.counters.clear() loss fn(x) # 1. 仅使用 torch.compile loss.backward(retain_graphTrue) assert(torch._dynamo.utils.counters[stats][unique_graphs] 3) torch._dynamo.utils.counters.clear() # 2. torch.compile compiled autograd with torch._dynamo.compiled_autograd.enable(torch.compile(backendaot_eager)): loss.backward() # 反向为单图 assert(torch._dynamo.utils.counters[stats][unique_graphs] 1)结果对比非常直观第一种情况仅torch.compile由于fn内有 2 次 graph breakbackward 也被切成了3 张图unique_graphs 3第二种情况叠加 Compiled Autograd尽管 forward 有断点backward 仍被完整追踪为1 张图unique_graphs 1。这里利用torch._dynamo.utils.counters[stats][unique_graphs]统计唯一图数量作为断言依据是一种非常实用的可观测手段。示例中显式指定backendaot_eager让读者聚焦图结构而忽略 Inductor 后端细节。注意即使启用 Compiled AutogradDynamo 在追踪由 Compiled Autograd 捕获的 backward hooks时仍然可能发生 graph break。也就是说hooks 被“捕获”不等于 hooks 一定“不产生断点”这一点在排查性能问题时需要牢记。场景二backward hooks 可以被捕获torch.compile(backendaot_eager) def fn(x): return x.sum() x torch.randn(10, 10, requires_gradTrue) x.register_hook(lambda grad: grad10) loss fn(x) with torch._dynamo.compiled_autograd.enable(torch.compile(backendaot_eager)): loss.backward()在纯 AOTAutograd 场景下x.register_hook(lambda grad: grad10)注册的钩子不会进入编译图而启用 Compiled Autograd 后图中会出现一个call_hook节点Dynamo 随后会把它内联展开为实际的钩子逻辑。下图即为日志中捕获到的call_hook节点红色框标注从日志节点的具体形态看call_hook节点会携带hook_type tensor_pre_hook等元信息并引用其输入张量如getitem_2、aot0_expand——这意味着钩子逻辑真正进入了图优化流程可以被后端进一步融合。两类常见的重编译recompilation原因Compiled Autograd 由于捕获的图更大、缓存键对图结构更敏感重编译比普通torch.compile更容易发生。以下是原文档给出的两类最常见诱因。原因一loss 的 autograd 结构发生变化torch._dynamo.config.compiled_autograd True x torch.randn(10, requires_gradTrue) for op in [torch.add, torch.sub, torch.mul, torch.div]: loss op(x, x).sum() torch.compile(lambda: loss.backward(), backendeager)()每次迭代调用不同的算子add/sub/mul/div导致loss追踪到不同的 autograd 历史即反向图中出现不同类型的Backward0节点。日志中会出现如下提示Cache miss due to new autograd node下图依次展示了四轮迭代中因新增GraphRoot、SubBackward0、MulBackward0、DivBackward0节点而触发的缓存未命中与重编译原因二张量形状发生动态变化torch._dynamo.config.compiled_autograd True for i in [10, 100, 10]: x torch.randn(i, i, requires_gradTrue) loss x.sum() torch.compile(lambda: loss.backward(), backendeager)()循环中x的形状在(10, 10)与(100, 100)之间变化。第一次形状变化后Compiled Autograd 会把x标记为动态形状张量此后每次形状变化都会触发重编译日志提示为Cache miss due to changed shapes从日志细节看重编译原因是GraphRoot与AccumulateGrad节点的形状发生变化框架会将相关维度的size idx标记为动态在重编译后的图中张量维度不再写死为具体数值而是以符号变量如s1、s2表示——这正是 PyTorch 动态形状dynamic shapes机制的体现。缓解思路对结构稳定的训练循环尽量保持每步计算的算子序列与张量形状一致避免无谓的缓存失效区分“必要重编译”输入形状确实改变与“非必要重编译”同形状反复编译必要时借助TORCH_LOGScompiled_autograd_verbose日志定位触发点对确实动态的形状可结合torch._dynamo.mark_dynamic等动态形状标注手段相关内容可参考 torch_export_tutorial.py 中关于动态形状Dim的讨论提前告知编译器减少首次运行时的编译抖动。总结与后续路径本文围绕 Compiled Autograd 完成了四件事定位问题解释了 AOTAutograd 部分捕获 backward 的两大限制前向断点传导、hooks 不入图以及 Compiled Autograd 通过直接集成 autograd 引擎捕获更大 backward 图的工作原理与自身代价上手使用演示了全局开关torch._dynamo.config.compiled_autograd True与上下文管理器torch._dynamo.compiled_autograd.enable(...)两种启用方式并拆解了从 Dynamo 捕获、AOTDispatcher 切分、Inductor 生成到 Compiled Autograd 引擎补全全图的完整调用链观测调试介绍了TORCH_LOGScompiled_autograd与TORCH_LOGScompiled_autograd_verbose两级日志、aot0_前缀的节点命名规律以及unique_graphs计数断言技巧排查重编译覆盖了“新增 autograd 节点”与“形状动态变化”两类最常见的缓存未命中原因及其日志特征。Compiled Autograd 是torch.compile生态中仍在快速演进的一环。想继续深入可以在本仓库中依次阅读 torch_compile 入门掌握 graph break、fullgraph等基础概念、torch_compile 全流程示例真实模型端到端实践并关注 dev-discuss 社区对 Compiled Autograd 的深入剖析文章。赞分享示例工程【免费下载链接】tutorialsPyTorch tutorials.项目地址https://gitcode.com/gh_mirrors/tuto/tutorials点击查看免费下载相关推荐leetcode 仓库实战精讲用 NumPy 手写多层反向传播Multi-Layer Backpropagation吃透 PyTorch autograd 的底层原理leetcode 仓库实战精讲用 NumPy 手写多层反向传播Multi Layer Backpropagation吃透 PyTorch autogra示例工程教程MXNet autograd 自动微分实战指南从梯度计算到自定义反向传播MXNet autograd 自动微分实战指南从梯度计算到自定义反向传播 本教程是 MXNet Crash Course 快速入门系列的第三步围绕 mxne深度学习人工智能机器学习分布式训练终极指南如何用DragonDiffusion实现5大核心图像编辑功能从入门到精通终极指南如何用DragonDiffusion实现5大核心图像编辑功能从入门到精通 DragonDiffusion是一款基于扩散模型的革命性图像编辑工具它能上一篇Million.js插件系统自定义编译规则与扩展功能下一篇SQLModel终极指南为什么它是Python数据库开发的最佳选择创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表