首页
/ PyTorch 高阶控制流算子 torch.map:批量映射的语义、导出与源码实现解析

PyTorch 高阶控制流算子 torch.map:批量映射的语义、导出与源码实现解析

2026-09-08 11:59:10作者:范垣楠Rhoda

本篇技术指南围绕 PyTorch 官方文档中的 Control Flow Map 算子展开,系统讲解 torch._higher_order_ops.maptorch.map)的结构化控制流语义、典型用法、面向 torch.export 的动态形状导出流程以及底层分发(dispatch)与实现机制。读完本文,你将掌握 map 的批量"按切片应用函数"编程模型、其三条使用限制及其在源码中的校验位置,并能结合源码理解它如何被转换为 torch.ops.higher_order.map_impl 调用与子图。

torch.map 是什么

torch.map 是 PyTorch 中一个结构化控制流算子(structured control flow operator,简称 HOP),它的作用是把一个函数应用到输入张量(或张量组成的 PyTree)的**首维(leading dimension)**上。也就是说,map 把输入沿第 0 维切分为若干切片,对每个切片独立执行一次函数体 f,最后把各步结果重新堆叠起来,作为输出整体返回。

官方文档给出了它的逻辑等价实现(见 map.md):

def map(
    f: Callable[[PyTree, ...], PyTree],
    xs: Union[PyTree, torch.Tensor],
    *args,
):
    out = []
    for idx in range(xs.size(0)):
        xs_sliced = xs.select(0, idx)
        out.append(f(xs_sliced, *args))
    return torch.stack(out)

需要强调的是,这只是帮助理解语义的伪代码:真实实现绝不会在 Python 层写这种串行循环,而是把 map 作为一个高阶算子注册到 dispatch 体系里,交由各后端(eager、FakeTensor、proxy tracing、functionalization 等)分别处理,这也是"结构化"三字的含义——算子本身保持独立节点,可以被 torch.export、torch.compile 等下游链路识别和改写。

从当前仓库源码看,map 的用户入口定义在 torch/_higher_order_ops/map.py#L111-L190,并已在 torch/_higher_order_ops/init.py 中通过 from torch._higher_order_ops.map import map 对外导出。

原型状态提示:官方文档与源码 docstring 均明确警告——torch._higher_order_ops.map 是 PyTorch 的原型(prototype)功能,目前仍可能遇到错误编译(miscompile),使用前应留意其功能分级。

核心用法示例:对 batch 逐切片应用函数

最典型的场景是把一个只接受单样本的函数 f 应用到一批样本上。官方文档给出如下示例(可完整运行):

import torch
from torch._higher_order_ops import map

def f(x):
    return x.sin() + x.cos()

xs = torch.randn(3, 4, 5)  # batch of 3 tensors, each 4x5
# Applies f to each of the 3 slices
result = map(f, xs)  # returns tensor of shape [3, 4, 5]
print(result)

这里 xs 形状为 [3, 4, 5],map 会把首维 3 视为 batch 维,f 被依次作用在每个 [4, 5] 的切片上,最终输出仍是 [3, 4, 5] 的张量——行为上与在 for 循环里对 xs[i] 逐个调用 f 等价,但 map 保留了结构化的算子形态,便于后续变换。

更丰富的函数签名

根据源码 docstring(map.py#L133-L155),f 的输入 x 既可以是单个张量,也可以是"张量的嵌套 dict、list";*args 是可选的额外输入,会被透传给 f 的每一步。例如:

def f(xs):
    return xs[0] + xs[1] + const1 + const2

xs = [torch.randn(2, 3), torch.randn(2, 3)]
const1 = torch.randn(2, 3)
const2 = torch.randn(2, 3)
# returns a tensor of shape [2, 2, 3]
torch._higher_order_ops.map(f, xs)

该例表明:xs 可以是张量列表,map 会同时按首维切分其中每个张量并保持结构;const1const2 这类不随 batch 变化的张量放在 *args 中,作为每一步共享的只读参数。此外,map 还能"自动判断读取依赖"——部分额外输入可以省略,由框架自动补全。

导出(Export)与动态形状:从 map 到 map_impl 子图

map 的一大价值在于可以进入 torch.export 的导出链路,把控制流与函数体固化成可进一步变换和部署的图。官方文档用一个支持可变 batch size 的动态形状示例演示了这一过程:

class MapModule(torch.nn.Module):
    def __init__(self):
        super().__init__()

    def forward(self, xs: torch.Tensor) -> torch.Tensor:
        def body_fn(x):
            return x.sin() + x.cos()

        return map(body_fn, xs)

mod = MapModule()
inp = torch.randn(3, 4)
ep = torch.export.export(mod, (inp,), dynamic_shapes={"xs": {0: torch.export.Dim.DYNAMIC}})
print(ep)

导出结果中会出现两个值得注意的现象(官方文档明确指出,源码亦可佐证):

  1. torch.maplower 为 torch.ops.higher_order.map_impl:真正注册进 dispatcher 的高阶算子是 MapImpl(HigherOrderOperator),其构造名即为 "map_impl"(见 map.py#L42-L44)。
  2. 函数体 body_fn 变成顶层图模块的一个子图属性(sub-graph attribute):在 torch.fx.experimental.proxy_tensor 的代理模式下追踪时,trace_map 会用 reenter_make_fx(f) 对首维切片(示例输入)单独构建一张 body 图,再通过 proxy_mode.tracer.root.register_module(next_name, body_graph) 把它挂到主图模块名下(默认命名形如 body_graph_0,见 map.py#L308-L341)。随后创建的 map_impl proxy 节点会以该子图模块为参数,从而形成"顶层图 + 内嵌 body 子图"的结构化表达。

导出时指定的 torch.export.Dim.DYNAMIC 会把 xs 第 0 维标记为动态维,因此在 body 子图的追踪与 FakeTensor 形状推导中,PyTorch 刻意避免"遍历 batch 维"这类会引入符号尺寸 guard 的操作,而是统一改用 first_slice_copy 取第一片来推断单步输出形状(见下文"源码级机制")。

单步输出的形状广播

在 FakeTensor 模式(map_fake_tensor_modemap.py#L364-L374)下,实现只需:

  1. first_slice_copyxs 每个张量的第一片;
  2. 执行 f 得到单步示例输出;
  3. 通过 _broadcast_to_batch 把单步输出的每个张量 unsqueeze(0).expand(batch_size, *shape)clone,从而无循环地构造出带 batch 维的假输出元数据。

这与 torch.stack 的结果布局保持一致(源码注释明确指出"Use contiguous_format to match torch.stack behavior")。

Restrictions:使用限制与其背后的源码校验

官方文档列出了 map 的三条限制,结合源码我们可以逐条找到对应的实现约束:

  1. 被映射的 xs 只能由张量组成 入口函数会对 xspytree.tree_flatten 后逐一检查,非张量会直接抛出 RuntimeError("Mapped xs can only consist of tensors. Got xs ...")map.py#L158-L161)。

  2. xs 中所有张量的首维必须一致且非零 源码先取第一个张量的首维 shapes[0][0] 作为统一 batch 大小,若为 0 则报 "Leading dimensions of mapped xs cannot be 0.";若任一其他张量的首维与之不同,则报 "Leading dimensions of mapped xs must be consistent. Got shapes ..."map.py#L163-L171)。该约束保证每个 batch 索引处都能取到结构对齐的切片。

  3. 函数体不得修改(mutate)输入 这条限制与 functionalization 流程强相关。在 map_functionalize 实现中(map.py#L377-L416),框架会对 body 函数做函数化(functionalize)包装,并通过 _check_alias_and_mutation(f, example_inputs, "map", pre_dispatch) 显式检查 body 是否存在别名与就地修改行为;同时追踪路径上只使用 first_slice_copy 构造的切片示例,避免真正的 batch 迭代。一旦 body 违反无副作用约束,检查会将其拒绝。

值得补充的是:源码中 MapImpl.gen_schema 的注释还说明了可变性契约的更精细语义map.py#L78-L88)——xs 的相邻切片存储互不相交(storage-disjoint),因此各迭代对 xs 的就地写彼此无竞态;但 pos_args(额外参数)中同一个张量会被每一轮共享,若 body 修改它就会让各迭代相互依赖、破坏"迭代彼此独立"的语义契约,在并行化 lower 时引入数据竞争。需要"共享可变缓冲区"的场景应改用 scan / while_loop(顺序迭代本就是它们的语义),相关文档见 docs/source/higher_order_ops 目录下的 scan.mdwhile_loop.md

源码级机制:map 的分发与各后端实现

理解 map 的底层原理,关键看 torch/_higher_order_ops/map.py 中围绕 map_impl 注册的多个实现。map 的完整调用链可概括为:

用户调用 map(f, xs, *args)
  → 扁平化 xs/args、做三条限制校验(map.py#L158-L171)
  → 构造 wrapped_fn 统一还原 PyTree 结构
  → map_impl(inner_f, flat_xs, flat_args)
      → 按当前 DispatchKey 分发给对应 impl

以下是各关键分发实现的职责:

分发实现 注册方式 核心职责
map_dense map_impl.py_impl(DispatchKey.CompositeExplicitAutograd)L344-L347 eager 参考语义:用 _unstack_pytree 沿首维拆开 xs,逐片调用 f,再用 _stack_pytree 堆叠结果,对应文档开头给出的伪代码
map_autograd map_impl.py_autograd_implL350-L354 包一层自定义 autograd 的 MapAutogradOp,在其 backward 中基于保存的前向输入重建反向图 bw_f,再次调用 map_impl 批量求梯度(L193-L284)。注意该能力同样处于原型阶段,文档 docstring 仍提示可能 miscompile
map_proxy_torch_dispatch_mode map_impl.py_impl(ProxyTorchDispatchMode)L357-L361 调用 trace_map,把 body 追踪为独立子图模块并产出 map_impl proxy 节点(即 export/fx 追踪路径)
map_fake_tensor_mode @register_fake(map_impl, ...)L364-L374 形状/元数据推导:用第一片示例输出广播出完整 batch 形状,不真正执行
map_functionalize map_impl.py_functionalize_implL377-L416 函数化语义:将 body 用 ctx.functionalize 包装,检查别名与修改,必要时走 do_auto_functionalize_v2 自动函数化路径

可以看出,文档中"用第一片代替全量迭代"的策略在 proxy 追踪与 FakeTensor 两条路径中被反复使用(源码多处注释:"Use first_slice_copy instead of _unstack_pytree to avoid iterating over batch dim, which would guard on symbolic sizes")。这是 map 得以在动态形状(batch 维为符号量)下仍可完成图构建的关键技巧:只要不遍历 batch 维,就不会产生基于符号尺寸的 guard。

这些实现的测试覆盖可参考 test/functorch/test_control_flow.pytest/dynamo/test_higher_order_ops.pytest/inductor/test_control_flow.py 等文件中对 map 行为、导出形态与编译路径的断言。

实践小结与注意事项

  • 导入路径from torch._higher_order_ops import map(在源码仓库环境内可用;对应算子为 torch.ops.higher_order.map_impl)。
  • 适用场景:希望对一批同构样本应用同一函数、且函数体无跨样本副作用时,map 能把这一过程表示为结构化控制流,便于导出(export)与后续变换。
  • batch 维约定:map 只作用于输入的首维,其余维度形状在各切片间保持一致;多张量输入时各首维必须相等且非零。
  • 副作用边界:函数体禁止就地修改输入;如需顺序迭代并共享可变状态,请改用 scan / while_loop。
  • 成熟度:本算子处于原型阶段,autograd 与编译组合路径可能仍存在缺陷;生产使用前应在目标用例上充分验证正确性,并留意功能分级公告。

结合官方文档 docs/source/higher_order_ops/map.md 与本仓库实现 torch/_higher_order_ops/map.py,开发者可以完整掌握 map 的语义模型、导出形态与底层分发机制,进而在自己的批量处理或模型导出流程中安全、正确地使用它。

热门项目推荐
相关项目推荐

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
docsdocs
暂无描述
Markdown
900
5.83 K
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.14 K
2.76 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
860
1.35 K
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
927
1.85 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.94 K
1.02 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
533
603
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.37 K
1.46 K
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
548
396
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.04 K
527