ARTICLE DETAIL

资讯详情

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

PyTorch 高阶算子 torch.scan 完全指南:结构化循环控制流、cumsum/cumprod 递归与 torch.export 导出实战

PyTorch 高阶算子 torch.scan 完全指南:结构化循环控制流、cumsum/cumprod 递归与 torch.export 导出实战 PyTorch 高阶算子 torch.scan 完全指南结构化循环控制流、cumsum/cumprod 递归与 torch.export 导出实战【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorchtorch.scan是 PyTorch 中一个结构化的控制流高阶算子Higher Order Operator用于执行带 combine 函数的包含式扫描inclusive scan可统一表达 cumsum、cumprod 等累积运算以及更一般的递归recurrence计算。本文基于当前仓库的官方文档 scan.md结合其核心实现 torch/_higher_order_ops/scan.py 与测试用例 test/functorch/test_control_flow.py系统讲解 scan 的语义、用法、导出流程、限制条件、内部实现与反向传播原理帮助你在实际项目中安全、高效地使用这一原型特性。Scan 是什么语义与结构化控制流定位torch.scan是一个结构化的控制流算子执行的是包含式扫描给定初始携带状态carryinit与输入序列xs沿着dim维逐个切片调用combine_fn(carry, x_slice)得到新的 carry 与当前步的输出y最后把每一步的y堆叠stack成输出张量。官方文档给出了它的逻辑等价实现scan.mddef scan( combine_fn: Callable[[PyTree, PyTree], tuple[PyTree, PyTree]], init: PyTree, xs: PyTree, *, dim: int 0, reverse: bool False, ) - tuple[PyTree, PyTree]: carry init ys [] for i in range(xs.size(dim)): x_slice xs.select(dim, i) carry, y combine_fn(carry, x_slice) ys.append(y) return carry, torch.stack(ys)可以看到scan 的语义非常直观init是初始 carry每一步combine_fn接收上一步的 carry 和当前切片返回(next_carry, y)最终返回最终 carry与沿扫描维堆叠的输出。值得注意的是这一伪代码只是为了表达语义。实际实现见下文并没有真正调用PyTree.select这样的接口而是把 pytree 展平后按叶子处理并且对输出做了预分配、对dim做了规范化、对reverse做了翻转处理。Prototype 警告torch.scan目前是 PyTorch 的原型prototype特性官方文档明确提示你可能会遇到 miscompiles编译错误/错误编译scan.md。在将其用于生产环境前请务必评估风险。快速上手用 scan 实现累积求和cumsum最小可用示例官方文档给出了一个用 scan 计算累积求和的示例scan.mdimport torch from torch._higher_order_ops import scan def add(carry: torch.Tensor, x: torch.Tensor): next_carry carry x y next_carry.clone() # clone to avoid output-output aliasing return next_carry, y init torch.zeros(1) xs torch.arange(5, dtypetorch.float32) final_carry, cumsum scan(add, initinit, xsxs) print(final_carry) print(cumsum)运行结果与实现注释中的说明一致见 scan.pyfinal_carry为tensor([10.])即01234cumsum为tensor([[0.], [1.], [3.], [6.], [10.]])是每一步的累积和切片堆叠而成。为什么要 clone示例中 combine 函数里出现了next_carry.clone()。这是 scan 的一条硬性约束combine_fn 的输出不能与任何输入存在别名关系aliasingscan.md。上面的add中next_carry与返回值y指向同一份存储如果不 clone就会构成 output-output 别名违反算子约束。同样地如果要把某个输入原样返回也必须 clone。在实现层面scan_functionalizescan.py会通过_check_alias_and_mutation对 combine_fn 的输入进行别名与可变性检查违反约束会直接报错。导出与部署把 scan 编进 TorchScript/ExportedProgramscan 的另一个重要应用场景是模型导出通过torch.export把包含 scan 的模块导出为 ExportedProgram便于后续的图变换与部署。官方文档给出了一个使用**动态形状dynamic shapes**支持可变序列长度的示例scan.mdclass ScanModule(torch.nn.Module): def forward(self, xs: torch.Tensor) - tuple[torch.Tensor, torch.Tensor]: def combine_fn(carry, x): next_carry carry x return next_carry, next_carry.clone() init torch.zeros_like(xs[0]) return scan(combine_fn, initinit, xsxs) mod ScanModule() inp torch.randn(5, 3) ep torch.export.export(mod, (inp,), dynamic_shapes{xs: {0: torch.export.Dim.DYNAMIC}}) print(ep)要点说明init torch.zeros_like(xs[0])保证了 init 与每步切片形状一致这里xs[0]形状为(3,)xs每步切片也是(3,)dynamic_shapes{xs: {0: torch.export.Dim.DYNAMIC}}将序列长度维第 0 维声明为动态维从而支持可变序列长度官方文档特别提示导出后combine 函数会变成顶层图模块的一个子图属性sub-graph attributescan.md。这与trace_scan中的实现吻合——它会为 combine 函数生成一个独立的 GraphModule 并注册为scan_combine_graphscan.py。完整 API 参考函数签名scan定义在 torch/_higher_order_ops/scan.py完整签名为def scan( combine_fn: Callable[ [pytree.PyTree, pytree.PyTree], tuple[pytree.PyTree, pytree.PyTree] ], init: pytree.PyTree, xs: pytree.PyTree, *, dim: int 0, reverse: bool False, length: int | None None, ) - tuple[pytree.PyTree, pytree.PyTree]:参数说明参数类型默认值说明combine_fnCallable必填二元函数类型为(Tensor, Tensor) - (Tensor, Tensor)当xs是 pytree 时为(pytree, pytree) - (pytree, pytree)。第一个输入是上一步/初始 carry第二个输入是沿dim的一个输入切片第一个输出是下一步 carry第二个输出是输出序列的一个切片。该函数必须是纯函数当前不支持 lifted 参数且不得有副作用initTensor 或张量叶子的 pytree必填初始扫描 carry其 pytree 结构必须与combine_fn第一个输出carry一致xsTensor、张量叶子的 pytree 或None必填输入张量或张量 pytree。当提供了length时可以传None此时combine_fn每步收到xNone计数器循环模式dimint0扫描的维度reverseboolFalse是否沿dim反向扫描lengthint 或 NoneNone可选扫描迭代次数。当xs含张量叶子时若给出则必须等于xs.shape[dim]仅作一致性校验当xs无叶子None或空 pytree时length决定迭代次数combine_fn每步收到xNone。length0且无 xs 张量仅在 eager 模式支持torch.compile下不支持返回值final_carry与init具有相同 pytree 结构的最终 carryout张量叶子组成的 pytree每个叶子是沿第 0 维堆叠的输出每个切片对应一次迭代的输出。若扫描维大小为 0则final_carry等于init不变各输出叶子沿dim大小为 0此时final_carry对init的梯度是单位映射identity而非零因为循环体从未被调用、carry 原样通过scan.py。限制Restrictions逐条解析官方文档列出了四条硬性限制scan.md这里结合源码逐一展开carry 元数据一致性combine_fn返回的next_carry必须与init具有相同的 shape 与 dtype。在trace_scan中会调用check_meta_consistency(init_fake_tensors, carry_fake_tensors, init, carry, ...)对 init 与 carry 的元数据做一致性校验scan.py并检测到不匹配时抛出错误。同时也要求 carry 的 pytree 结构与 init 一致。禁止原地修改输入combine_fn不得 in-place 修改其输入如需修改必须先 clone。在scan_functionalize中_check_alias_and_mutation会对 combine_fn 的输入做可变性检查scan.py。实现注释还说明这一禁止输入变异限制计划在推理场景下很快解除scan.py。不得修改函数外部创建的 Python 变量如 list/dictcombine_fn 必须是纯函数不得有副作用scan.py。输出不得与任何输入别名combine_fn 的输出不能是输入的视图或原样返回需要 clone。这与gen_schema中描述的变异语义一致init 与 xs 均不可变——init 只是初始 carry对 init 的 in-place 更新只会影响第 0 步xs 的每个切片在存储上互不相交对xs[t]的修改无法被第 t1 步观察到scan.py。给开发者的建议如果你确实需要在循环过程中维护类似 xs 的 in-place 更新实现注释给出的路径是把该缓冲区作为additional_inputs传入并在 combine_fn 内部对它做索引写入scan.py。深入实现scan 在源码中是如何工作的入口流程展平、校验、维度规范化scan()入口scan.py做了以下关键步骤pytree 展平把init与xs分别展平为叶子列表与结构 spec保证 combine_fn 输入顺序与输出顺序一致scan.pylength 短路逻辑当length给出且 xs 无张量叶子时会用torch.zeros(length, dtypetorch.int64)伪造一个长度为length的迭代计数器包装后的 combine_fn 丢弃切片并把xNone传给用户函数scan.pylength0且无 xs 张量时通过探测combine_fn(init, None)输出结构来构造空输出_build_empty_output_for_length_zeroscan.py参数校验_validate_input检查 combine_fn 可调用、dim为 int、reverse为 bool、init 叶子均为 Tensor、xs 叶子均为 Tensor 且每个 xs 叶子ndim dim、各 xs 叶子扫描维大小一致以及length与xs.shape[dim]的一致性scan.py维度规范化与重排用utils.canonicalize_dim规范化dim支持负索引再通过torch.movedim(elem, dim, 0)把扫描维移到 0统一在 dim 0 上扫描reverseTrue时先torch.flip输入scan.py包装 combine_fn通过wrap_combine_fn_flat把用户函数包装为展平叶子版本保证调用约定scan.py 与 scan.py调用与还原经_maybe_compile_and_run_fn分发执行结果在reverseTrue时再 flip 回来dim ! 0时再movedim回原维度scan.py。底层算子与 eager 实现scan的底层是ScanOp(HigherOrderOperator)scan.py其 eager 实现scan_op_dense挂在CompositeExplicitAutogradkey 上实际调用generic_scanscan.py。generic_scanscan.py是真正的循环实现值得注意的优化点用first_slice_copy取第 0 个切片先跑一次 combine_fn既完成输出形状推断用于预分配又直接产出第 0 个真实结果避免了对有副作用的算子多调用一次历史实现曾多调用一次导致对副作用算子不正确见 scan.py 的注释输出张量outs按[num_elems] list(e.size())预分配配合idxs索引矩阵与scatter_把每步输出写入对应位置scan.py每步通过elem.select(dim, i)取第 i 个切片scan.py。各种 dispatch 实现除了 eagerscan 还实现了多个 dispatch 后端ProxyTorchDispatchModetracetrace_scanscan.py在关闭 proxy 追踪的模式下用reenter_make_fx把 combine_fn 转成 FX GraphModule校验 init/carry 元数据一致用stack_y构造堆叠输出并把 combine 图注册为顶层模块的子模块scan_combine_graphFakeTensorModescan_fake_tensor_modescan.py仅对第 0 个切片跑一次 combine_fn 推断输出元数据再用stack_y广播出完整形状用于 shape 推断Functionalizescan_functionalizescan.py负责别名/变异检查可自动功能化auto-functionalize以支持带副作用的 combine_fnVmap批处理规则scan_batch_rulescan.py先把批维度移动到末维用_VmapCombineFnWrapper包装 combine_fn 使其兼容批处理再调用 scan_op最后把批维度还原。参考实现_fake_scanscan.py 还提供了_fake_scan——一个纯 Python 的参考实现仅用于测试。测试代码大量用它作为期望结果与真实scan对比如 test/functorch/test_control_flow.py。阅读_fake_scan是理解 scan 语义含 reverse、length、movedim 行为最直接的方式。反向传播scan 的 autograd 实现scan 的反向传播是一个值得深入的点。它定义在ScanAutogradOp(torch.autograd.Function)与ScanAutogradImpl中scan.py。核心思路是scan 的反向也是一个 scan反向的扫描前向通过HopGraphMinCutPartitioner把 combine_fn 切分成前向子图fw_gm与反向子图bw_gmscan.py_optimize_forward_intermediates把前向中间量按来源分为 4 类并制定处理策略scan.py 中的枚举KEEP真正的中间张量随每步变化作为 ys 保存CLONE是 carryinit的一部分每步会被新值替换需 clone 后堆叠保存REMOVE_XS是 xs 的一部分xs 只读可直接保存原始输入供反向使用REMOVE_ADDITIONAL_INPUTS是 additional_inputs 的一部分同样只读直接保存。call_backwardscan.py构造反向 scanbw_init (grad_carry, grad_additional_inputs)bw_xs (fw_intermediates, grad_ys)然后以reverseTrue再跑一次 scan其中grad_additional_inputs在每步用加法累积grad_carry逐迭代传递grad_x作为 ys 输出并在结束后堆叠翻转回原方向。此外_break_bw_input_output_aliasingscan.py会检测反向子图输出是否与输入占位符别名如直接返回不 requiring grad 的输入对应的zeros_like并对这类输出插入 clone以满足 scan_op 的无别名不变量、避免 dynamo 下的UncapturedHigherOrderOpError。测试验证scan 在仓库中的覆盖情况torch.scan有完整的测试覆盖。主要测试位于 test/functorch/test_control_flow.py其中get_scan_combine_fntest/functorch/test_control_flow.py提供了多种 combine 函数模板add/mul/div点对点运算、S5 状态空间模型算子s5_operator、different_input_size_operator不同输入/输出尺寸、tuple_fct元组 pytree、complex_pointwise字典列表元组混合嵌套 pytree、RNN带参数的隐状态递归等test_scan_y_less_ndim_then_dimtest/functorch/test_control_flow.py等用例覆盖 y 输出维度小于扫描维的边界情况test_scan_compiletest/functorch/test_control_flow.py在 eager 与 torch.compile 多种编译模式下用_fake_scan作为期望结果对比scan的输出。这些测试是验证你的 combine_fn 是否符合 scan 约束的最佳参考如果你不确定某个 combine_fn 能否被 scan 接受可以先看看测试中是否有相似模式的先例。实战要点总结记住三不约束不改输入in-place、不别名输出 clone、不改外部变量纯函数。这是 scan 使用中 90% 报错来源init 与 next_carry 必须同构同型pytree 结构、shape、dtype 都要一致默认dim0非 0 维度扫描、负索引维度都支持内部会 canonicalize 并 movedim 到 0reverseTrue实现反向扫描可用于双向 RNN 等场景无 xs 的计数器循环通过lengthNxsNone实现纯迭代次数驱动的循环combine_fn每步收到xNone导出用torch.export 动态形状combine_fn 会变成导出图的一个子图属性可继续走图变换与部署流程当前为 prototype可能遇到 miscompiles生产使用前务必充分测试并关注 PyTorch 特性分类feature classification的更新。如需进一步研究实现细节建议精读 torch/_higher_order_ops/scan.py入口/校验/各 dispatch 实现、test/functorch/test_control_flow.py丰富测试用例以及torch._higher_order_ops.utils中的first_slice_copy、check_meta_consistency等辅助函数。【免费下载链接】pytorchTensors and Dynamic neural networks in Python with strong GPU acceleration项目地址: https://gitcode.com/GitHub_Trending/py/pytorch创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表