ARTICLE DETAIL

资讯详情

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

PyTorch torch.export 编程模型深度指南:从追踪原理到静态/动态语义实战

PyTorch torch.export 编程模型深度指南:从追踪原理到静态/动态语义实战 PyTorch torch.export 编程模型深度指南从追踪原理到静态/动态语义实战【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchtorch.export.export是 PyTorch 面向编译器栈TorchInductor、AOT 推理等的前端 API它以示例输入为样本将nn.Module静态追踪trace为一张仅含 ATen 算子的函数式计算图ExportedProgram。本文以 docs/source/user_guide/torch_compiler/export/programming_model.md 为骨架系统讲解 export 的追踪模式、静态/动态值的判定规则、输入类型约束、控制流处理、符号形状传播与模块状态访问规则并逐一给出仓库源码与可运行示例帮助你建立对torch.export.export处理代码方式的直觉写出一次导出、处处可运行的正确模型。Basics of Tracingexport 如何捕获一张图torch.export.export(mod, args)通过在示例输入example inputs上执行你的模型来捕获计算图它记录沿实际执行路径观察到的 PyTorch 算子与条件分支得到的结果是一张带元数据的单一 PyTorch 算子图只要后续输入满足同样的条件这张图就可以在不同的输入上运行。这张图的精确格式IR由 export IR 规范 定义。从仓库入口看export的函数签名位于 torch/export/init.py核心参数包括mod被追踪forward方法的模块args/kwargs示例输入dynamic_shapes动态形状规格Dim、dict/tuple 规格或新的ShapesSpec/ParamsSpecAPIstrict严格/非严格追踪开关preserve_module_call_signature为unflatten保留子模块调用约定的路径列表。导出成功后返回ExportedProgram定义见 torch/export/exported_program.py它可用torch.export.save/torch.export.load序列化详见 docs/source/user_guide/torch_compiler/export.md。Strict vs. Non-Strict 追踪torch.export.export提供两种追踪模式非严格模式non-strict用普通 Python 解释器执行程序代码行为与 eager 模式完全一致唯一区别是所有 Tensor 被替换为fake Tensor有形状等元数据但没有数据并包裹在Proxy 对象中从而把对其执行的所有算子记录进图同时捕获对 Tensor 形状的条件guard以保证生成代码的正确性。严格模式strict先用 TorchDynamoPython 字节码分析引擎对程序做符号分析而不真正执行代码再基于分析结果构图。这带来了超出形状条件的 Python 级安全性保证但代价是并非所有 Python 特性都被 TorchDynamo 支持。值得注意的事实是虽然文档写作时严格模式仍是默认强烈推荐使用非严格模式并将在未来成为默认因为对大多数模型而言形状条件已足以保证健全性soundness而 TorchDynamo 不支持某些 Python 特性反而是不必要的风险。在本仓库中这一推荐已落地为代码事实torch.export.export的strict参数默认值就是False见 torch/export/init.py且文档也说明切换该参数不会改变最终 IR 的规格与序列化方式。本文后续统一假设以非严格模式追踪即所有 Python 特性都被支持。理解核心概念静态值Static与动态值Dynamic理解torch.export.export行为的关键是区分静态值与动态值。静态值导出时固定并烧录进图静态值在导出时刻被固定在导出程序的多次执行之间不可改变。追踪中一旦遇到静态值就将其视为常量直接硬编码进图。当一个算子如x y的所有输入都是静态值时其输出被直接常量折叠constant-folded进图算子本身不会出现在图中。一个值被硬编码进图后我们称图对该值特化specialized了。原文档给出如下示例import torch class MyMod(torch.nn.Module): def forward(self, x, y): z y 7 return x z m torch.export.export(MyMod(), (torch.randn(1), 3)) print(m.graph_module.code) def forward(self, arg0_1, arg1_1): add torch.ops.aten.add.Tensor(arg0_1, 10); arg0_1 None return (add,) 这里y的示例值是3被当作静态值处理与7相加后把常量10烧录进图y 7这步在图中完全消失。动态值每次运行都可不同动态值在每次运行之间可以变化行为与普通函数参数一致传入不同输入函数都能给出正确结果。什么值静态、什么值动态判定规则取决于值的类型TensorTensor 的数据是动态的Tensor 的形状可由系统决定为静态或动态——默认所有输入 Tensor 的形状都是静态的用户可通过dynamic_shapes为任意输入 Tensor 指定动态形状属于模块状态的参数parameters与缓冲区buffers的 Tensor 形状永远是静态的其他 Tensor 元数据如device、dtype是静态的。Python 原生类型int、float、bool、str、None是静态的存在对应的动态变体SymInt、SymFloat、SymBool一般用户无需直接接触整数输入也可通过dynamic_shapes指定为动态。Python 标准容器list、tuple、dict、namedtuple结构list/tuple的长度、dict/namedtuple的键序列是静态的包含的元素递归套用上述规则即 PyTree 方案叶子类型为 Tensor 或原生类型。其他类包括 dataclass可通过 PyTree 注册见下文并遵循与标准容器相同的规则。输入类型默认支持与自定义 PyTree输入按上述规则被判定为静态或动态静态输入会被硬编码进图运行时传入不同值会报错这类输入多为原生类型值动态输入行为与普通函数输入一致这类输入多为 Tensor 值。默认允许的输入类型为Tensor、Python 原生类型int、float、bool、str、None、Python 标准容器list、tuple、dict、namedtuple。自定义输入类型注册 PyTree想用自定义类作为输入类型需要把该类注册为 PyTree。原文档示例用register_dataclass注册一个 dataclassdataclass class Input: f: torch.Tensor p: torch.Tensor import torch.utils._pytree as pytree pytree.register_dataclass(Input) class M(torch.nn.Module): def forward(self, x: Input): return x.f 1 torch.export.export(M(), (Input(ftorch.ones(10, 4), ptorch.zeros(10, 4)),))从源码看register_dataclass(cls, *, field_namesNone, drop_field_namesNone, serialized_type_nameNone)定义于 torch/utils/_pytree.py它是对更底层register_pytree_node的简化封装要求目标类具有 dataclass 语义torch.export.export的 docstring 也明确把已注册的 dataclass列为可接受输入类型之一见 torch/export/init.py。自定义类型注册后其字段会被递归展开为叶子Tensor 或原生类型从而获得与标准容器一致的静态/动态判定。可选输入未传入即被特化到默认值对于未传入的可选输入默认参数torch.export.export会将其特化到默认值。结果是导出程序要求调用方显式传入全部参数丢失默认值行为。原文档示例class M(torch.nn.Module): def forward(self, x, yNone): if y is not None: return y * x return x x # Optional input is passed in ep torch.export.export(M(), (torch.randn(3, 3), torch.randn(3, 3))) print(ep) ExportedProgram: class GraphModule(torch.nn.Module): def forward(self, x: f32[3, 3], y: f32[3, 3]): # File: /data/users/angelayi/pytorch/moo.py:15 in forward, code: return y * x mul: f32[3, 3] torch.ops.aten.mul.Tensor(y, x); y x None return (mul,) # Optional input is not passed in ep torch.export.export(M(), (torch.randn(3, 3),)) print(ep) ExportedProgram: class GraphModule(torch.nn.Module): def forward(self, x: f32[3, 3], y): # File: /data/users/angelayi/pytorch/moo.py:16 in forward, code: return x x add: f32[3, 3] torch.ops.aten.add.Tensor(x, x); x None return (add,) 对比两次导出可见传入y时图中保留了mul分支不传y时y被特化为None图中只剩add分支且签名中y不再有类型标注已被烧录。控制流静态分支与动态分支的处理torch.export.export支持控制流具体行为取决于被分支的值是静态还是动态。静态控制流透明支持基于静态值的 Python 控制流透明支持静态形状上的控制流也属于此类。静态值被烧录因此导出图永远不会看到静态值上的控制流if语句继续追踪导出时刻实际走到的那个分支for/while语句通过展开循环继续追踪。动态控制流区分形状依赖与数据依赖当控制流涉及的值是动态的它可能依赖动态形状或动态数据。由于编译器追踪时掌握的是形状信息而非数据二者对编程模型的影响截然不同。动态形状依赖Shape-Dependent的控制流当控制流涉及动态形状时多数情况下追踪期间我们同样知道该动态形状的具体值详见下文符号形状一节。此时称控制流是形状依赖的用动态形状的具体值把条件求值为True或False继续追踪如上文所述同时发射一个与刚求值条件对应的 guard。否则控制流被视为数据依赖我们无法把条件求值为True或False无法继续追踪必须在导出时报错。动态数据依赖Data-Dependent的控制流基于动态值的数据依赖控制流是被支持的但必须使用 PyTorch 的显式算子来继续追踪。直接用 Python 控制流语句基于动态值分支是不被允许的——编译器无法求值继续追踪所需的条件必须在导出时报错。PyTorch 为此提供了表达动态值上一般条件与循环的算子典型如torch.cond、torch.map。注意只有当你确实需要数据依赖控制流时才需要用它们。原文档给出了把数据依赖if改写为torch.cond的示例。x.sum() 0依赖输入数据改写后无需二选一追踪分支而是两个分支都被追踪class M_old(torch.nn.Module): def forward(self, x): if x.sum() 0: return x.sin() else: return x.cos() class M_new(torch.nn.Module): def forward(self, x): return torch.cond( predx.sum() 0, true_fnlambda x: x.sin(), false_fnlambda x: x.cos(), operands(x,), )从源码看torch.cond(pred, true_fn, false_fn, operands())定义于 torch/_higher_order_ops/cond.pypred是布尔表达式或单元素 Tensortrue_fn/false_fn是两个分支可调用对象operands是传递给分支的输入它作为高阶算子进入图两个分支都被保留供下游编译器在运行时根据pred选择执行。torch.map则定义于 torch/_higher_order_ops/map.py其语义等价于对xs的第一个维度做循环f并torch.stack结果用于表达动态循环。数据依赖控制流的一个特例是涉及数据依赖动态形状unbacked SymInt典型如依赖输入数据的中间 Tensor 形状。此时可以不使用控制流算子而是提供一条断言来决定条件为True还是False有了断言即可继续追踪并发射 guard。对应算子如torch._check仅当存在对数据依赖动态形状的控制流时才需要使用。原文档示例nz x.nonzero()的输出形状依赖输入数据用torch._check断言后即可继续追踪class M_old(torch.nn.Module): def forward(self, x): nz x.nonzero() if nz.shape[0] 0: return x.sin() else: return x.cos() class M_new(torch.nn.Module): def forward(self, x): nz x.nonzero() torch._check(nz.shape[0] 0) if nz.shape[0] 0: return x.sin() else: return x.cos()从源码看torch._check的底层实现在 torch/fx/experimental/symbolic_shapes.py 中大量用于约束 SymInt 之间的关系例如ShapeEnv内部通过torch._check(sym_int existing_symint)把新符号绑定到既有符号当你触发GuardOnDataDependentSymNode错误时错误信息还会直接给出torch._check(...)形式的修复建议可复制进代码。符号形状Symbolic Shapes基础追踪期间动态 Tensor 形状及其上的条件被编码为符号表达式静态形状及其条件则只是int与bool。符号symbol像一个变量描述一个动态 Tensor 形状。随着追踪推进中间 Tensor 的形状可能由更一般的表达式描述通常涉及整数算术运算——因为对大多数 PyTorch 算子输出形状可描述为输入形状的函数例如torch.cat输出的形状是其输入形状之和。同时遇到程序中的控制流时会创建布尔表达式通常涉及关系运算来描述被追踪路径上的条件这些表达式被求值以决定追踪哪条路径并记录在**形状环境shape environment**中用于守护被追踪路径的正确性以及求值后续创建的表达式。形状环境的实现即ShapeEnv类见 torch/fx/experimental/symbolic_shapes.py其evaluate_expr同文件 L8703 附近负责对符号布尔表达式求值并注册 guard。算子的 FakeMeta实现追踪期间程序以没有数据的 fake Tensor 执行因此一般无法调用 PyTorch 算子的真实实现——每个算子必须额外提供一个 fake即 meta实现输入输出都是 fake Tensor在形状等元数据行为上与真实实现一致。例如torch.index_select的 fake 实现用输入形状计算输出形状忽略输入数据、返回空数据def meta_index_select(self, dim, index): result_size list(self.size()) if self.dim() 0: result_size[dim] index.numel() return self.new_empty(result_size)Fake 实现的注册与调度相关代码可在 torch/fx/experimental/proxy_tensor.py代理张量机制与各算子的 meta 定义中找到。形状传播Backed 与 Unbacked 动态形状形状通过算子 fake 实现传播。理解动态形状传播的关键概念是backed与unbacked动态形状前者的具体值我们知道后者的具体值我们不知道。传播过程如下输入 Tensor 的形状可为静态或动态动态时由符号描述由于导出时用户提供了真实示例输入这些符号是 backed 的我们知道它们的具体值。算子输出形状由 fake 实现计算可为静态或动态动态时一般由符号表达式描述。进一步地若输出形状只依赖输入形状则当输入形状全为静态或 backed 动态时输出形状是静态的或 backed 动态的若输出形状依赖输入数据则它必然是动态的且因为我们无法知道其具体值它是 unbacked 的。控制流Guard 与断言遇到形状条件时若只涉及静态形状它是bool若涉及动态形状它是符号布尔表达式。对后者只涉及 backed 动态形状可用其具体值把条件求值为True/False随后向形状环境添加一条 guard声明对应符号布尔表达式为真/假继续追踪。涉及 unbacked 动态形状通常无法在不借助额外信息的情况下求值因此无法继续追踪必须在导出时报错用户应改用显式 PyTorch 算子如torch._check继续追踪。该信息会作为 guard 加入形状环境还可能帮助把后续遇到的其他条件求值为真/假。导出完成后backed 动态形状上的 guard 可理解为对输入动态形状的条件。它们会与导出时提供的动态形状规格dynamic shape specification核对——该规格描述了示例输入以及未来所有输入都需满足的动态形状条件。更精确地说动态形状规格必须在逻辑上蕴含生成的 guard否则导出时报错并附带对动态形状规格的修改建议反之若没有生成任何 backed 动态形状上的 guard尤其所有形状都是静态时则无需提供动态形状规格。通常动态形状规格会被转换为生成代码输入上的运行时断言。unbacked 动态形状上的 guard 会被转换为内联运行时断言插入在生成代码中该 unbacked 动态形状被创建的位置——典型场景是紧跟在数据依赖算子调用之后。允许的 PyTorch 算子与自定义算子所有 PyTorch 算子都被允许出现在导出图中。此外你还可以定义并使用自定义算子custom operators定义自定义算子时必须像其他 PyTorch 算子一样为其定义 fake 实现。原文档示例一个包装 NumPy 的自定义sin算子及其平凡的fake 实现torch.library.custom_op(mylib::sin, mutates_args()) def sin(x: Tensor) - Tensor: x_np x.numpy() y_np np.sin(x_np) return torch.from_numpy(y_np) torch.library.register_fake(mylib::sin) def _(x: Tensor) - Tensor: return torch.empty_like(x)torch.library.register_fake正是算子 fake 实现的注册入口见 torch/library.py。有时自定义算子的 fake 实现会涉及数据依赖形状。例如一个自定义nonzero的 fake 实现... torch.library.register_fake(mylib::custom_nonzero) def _(x): nnz torch.library.get_ctx().new_dynamic_size() shape [nnz, x.dim()] return x.new_empty(shape, dtypetorch.int64)这里new_dynamic_size()从算子注册上下文申请一个新的 unbacked 动态尺寸符号用于描述依赖输入数据的输出维度——这正是unbacked 动态形状在自定义算子侧的产生途径与上文中torch.nonzero等原生算子行为一致。模块状态读取Reads与更新Updates模块状态包括参数parameters、缓冲区buffers与普通属性regular attributes普通属性可以是任意类型参数和缓冲区永远是 Tensor。模块状态基于上述类型规则判定静态或动态例如self.training是bool因而是静态的任何参数或缓冲区都是动态的。模块状态中任何 Tensor 的形状都不能是动态的——这些形状在导出时刻固定不能在导出程序的多次执行之间改变。访问规则所有模块状态必须已初始化访问未初始化的模块状态会在导出时报错。读取模块状态总是允许的。更新模块状态需遵循以下规则静态普通属性如原生类型可以更新。读取与更新可自由交错且读取始终看到最近更新的值由于属性是静态的其值会被烧录生成代码中不会有任何实际 get/set 该属性的指令。动态普通属性如 Tensor 类型不可以更新若需更新必须在模块初始化时把它注册为缓冲区。缓冲区可以更新可以是就地更新如self.buffer[:] ...或整体赋值如self.buffer ...。参数不可以更新。参数通常在训练而非推理时更新建议在导出时使用torch.no_grad避免导出期间的参数更新。Functionalization函数化的效果被读取/更新的动态模块状态会被相应地提升lift为生成代码的输入/输出。导出程序在生成代码之外还保存参数与缓冲区的初始值以及其他 Tensor 属性的常量值。这意味着导出的ExportedProgram是自包含的下游运行器可以从初始值出发把被提升为输出的状态更新持久化实现真正的 ahead-of-time 函数式执行。总结一份面向实践的导出决策清单结合全文导出模型时可依此检查选择追踪模式默认strictFalse非严格享受完整的 Python 特性支持仅当需要额外的 Python 级安全保证时才开strictTrueTorchDynamo 符号分析Python 特性覆盖有限。盘点静态/动态值Tensor 数据动态、形状默认静态可用dynamic_shapes放宽原生类型静态容器结构静态、元素递归判定模块状态形状永远静态。处理输入默认类型开箱即用自定义类用pytree.register_dataclass注册可选参数会被特化导出后需显式传参。处理控制流静态值上的if/for/while透明支持分支追踪/循环展开数据依赖控制流改用torch.cond/torch.map数据依赖动态形状上的分支用torch._check断言。处理算子所有 PyTorch 算子可用自定义算子必须注册 fake 实现涉及数据依赖形状时用get_ctx().new_dynamic_size()声明 unbacked 尺寸。管理模块状态只读最安全静态属性可自由更新动态 Tensor 属性须注册为 buffer 才能更新参数不可更新导出时用torch.no_grad。依照这套规则写出的模型可被torch.export.export干净地捕获为一张带完整 guard/断言元数据的函数式图再交由torch.export.save/load序列化部署或交给下游编译器做进一步的推理优化。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表