PyTorch 原生并行:从 Tensor 到 DTensor
Quick Start
|
|
|
|
从 Tensor 到 DTensor
Tensor
普通 torch.Tensor 可以先理解成「一段数据 + 一组元信息」:
- 数据本体存放在某个 device 上,例如 CPU 或一张 GPU。
- 元信息描述这段数据应该如何被解释,例如
shape、dtype、stride、requires_grad。 - 算子看到的是一个单设备张量,所有输入、输出和中间结果默认也都落在这个单设备语义里。
PyTorch 内部会用类似 TensorMetadata 的结构描述一个 Tensor 的逻辑属性:
|
|
当 Tensor 扩展到多设备时,仅有这些元信息是不够的。系统还必须回答几个分布式问题:
- 这组设备如何组织?例如是一维数据并行 mesh,还是二维
dp x tpmesh。 - 全局 Tensor 如何映射到每个 rank?例如按第 0 维切分、每张卡保存一份副本,或每张卡只保存部分归约结果。
- 算子执行后新的分布式布局是什么?哪些算子可以纯本地计算,哪些算子需要插入
all-gather、all-reduce、reduce-scatter等通信。 - 用户代码应该看到全局 Tensor 语义,还是每个 rank 的局部 Tensor 语义。
DTensor
DTensor 的核心目标是把「多设备分布式张量」包装成一个仍然像 torch.Tensor 一样使用的对象:
|
|
其中最重要的是区分两层视角:
- Global view:用户代码看到的逻辑 Tensor,
shape、dtype、算子语义都按完整 Tensor 理解。 - Local view:每个 rank 实际持有的
_local_tensor,它可能只是 global Tensor 的一个 shard,也可能是完整副本,或者是尚未归约完成的 partial 结果。
因此,一个 DTensor 不只是多了几张卡上的存储,它还显式记录了分布式布局:
| 概念 | 作用 |
|---|---|
DeviceMesh |
描述参与这个 DTensor 的设备拓扑,以及每个 mesh 维度对应的通信组 |
Placement |
描述 Tensor 在每个 mesh 维度上的放置方式,例如 Shard、Replicate、Partial |
DTensorSpec |
把 DeviceMesh、Placement 和 Tensor 元信息组合起来,完整描述 DTensor 的布局 |
从编程模型看,DTensor 提供的是 SPMD 语义:所有 rank 运行同一份 Python 代码,但每个 rank 根据相同的 DTensorSpec 操作自己的 local shard。算子分发时,DTensor 会基于输入布局推导输出布局,并在必要时自动插入集合通信,让用户尽量按照单机 Tensor 的方式写代码。
DeviceMesh - Describing Device Topology
DeviceMesh 提供表达一组 device 布局的抽象,可以用一个多维数组表达,同时也提供 Mesh 内 device 通信的支持。
可以通过 init_device_mesh 来初始化一个 DeviceMesh:
|
|
对应可视化
|
|
|
|
对应可视化:
|
|
访问 sub-meshes:
|
|
Placement - Describing Tensor Distribution
Placements describe how a tensor is distributed across each dimension of the DeviceMesh.
一个 DTensor 的 placements 数量必须和 DeviceMesh 的维度数一致。例如一维 mesh 只需要一个 placement,二维 mesh 则需要两个 placement:
|
|
可以把 placements 理解成对每个 mesh 维度的逐项说明:
| Mesh 维度 | Placement | 含义 |
|---|---|---|
dp |
Replicate() |
沿数据并行维度复制完整 Tensor |
tp |
Shard(1) |
沿张量并行维度切分 Tensor 的第 1 维 |
DTensor 支持三类基础 placement:Shard、Replicate 和 Partial。
Shard(dim)
Shard(dim) 表示沿 Tensor 的第 dim 个维度切分数据,并把切分后的 shards 分发到 mesh 维度上的不同 rank。
|
|
逻辑上,用户仍然看到一个完整的 [8, 4] global Tensor:
|
|
如果在 4 张 GPU 上沿第 0 维切分,每个 rank 只持有 [2, 4] 的 local shard:
|
|
Shard 是节省显存的主要手段:global Tensor 的总数据量被拆到多个 rank 上。但它也意味着某些算子需要通信。例如,当一个算子要求完整 Tensor 时,DTensor 可能需要先 all-gather;当输出仍可保持分片时,则可以只在 local shard 上计算。
Replicate()
Replicate() 表示 mesh 维度上的每个 rank 都保存一份完整 Tensor 副本。
|
|
复制后的每个 rank 都持有完整 [8, 4]:
|
|
Replicate 的优点是读访问简单,很多 elementwise 算子可以直接在每个 rank 上独立执行;缺点是显存占用随副本数线性增加。数据并行里的参数副本、或者需要广播到所有 rank 的小张量,通常适合用 Replicate 表达。
Partial(reduce_op)
Partial(reduce_op) 表示每个 rank 持有的是某个 global Tensor 的「部分结果」,这些 partial values 需要通过归约操作才能变成完整语义的 DTensor。
|
|
常见的 reduce_op 包括 "sum"、"avg"、"product"、"max"、"min"。其中最常见的是 "sum",例如矩阵乘法或线性层在某个维度被切分后,每个 rank 只计算了输出的一部分累加项:
|
|
Partial 通常是算子传播过程中的中间布局,而不是用户手动构造数据时最常用的布局。它的核心意义是把「尚未完成归约」这件事显式记录在 DTensor 的 layout 中,这样后续算子可以决定是继续保留 partial 状态,还是在需要完整值时触发 all-reduce。
Multi-Dimensional Placements
对于多维 DeviceMesh,placements 是一个 tuple/list,每一项对应一个 mesh 维度。下面的例子使用二维 mesh:
|
|
含义是:
- 沿
dp维度使用Replicate():两组数据并行 rank 拥有相同的逻辑数据。 - 沿
tp维度使用Shard(1):每个dp组内部再按 Tensor 的第 1 维做列切分。
可视化如下:
|
|
这里同一列位置上的两个 dp rank 互为 replica,而同一行里的四个 tp rank 共同组成一个按列切分的 global Tensor。
作为对比, PyTorch 的 DTensor 和 OneFlow1 以及 GSPMD2 定义的区别:
| PT-D DistributedTensor | OneFlow’s SBP | GSPMD’s tensor sharding |
|---|---|---|
| Shard | Split | Tiled |
| Replicate | Broadcast | Replicated |
| Partial | Partial | Partially tiled = tiled + replicated |
DTensorSpec
DTensorSpec 完全表达了一个 DTensor 的元信息,分别由以下三个部分组成:
DeviceMesh对象来表达 DTensor 的 mesh 信息,Tuple[Placement]来表达 placements 方法TensorMetadata对象来表达 global tensor 的 meta 信息
|
|
实际举例:
|
|
DTensor
Torch 的 DTensor 在 torch.Tensor 类型上进行了简单的封装:
|
|
主要包括 _local_tensor 和 _spec:
_local_tensor是实际存储的torch.Tensor变量 (per rank)。_spec中存储了 DTensor 的全部元信息,对应于DTensorSpec字段- 包括 DeviceMesh、切分策略(Placements)以及传统 tensor 的属性信息,例如 shape、dtype 等
Creating DTensor
创建 DTensor 最常见有两条路径:
distribute_tensor():从一个 global logical tensor 出发,由 DTensor 负责 scatter / broadcast 到各个 rank。DTensor.from_local():从每个 rank 已经存在的 local tensor 出发,告诉 DTensor 这些 local tensor 共同组成什么 global layout。
两者的区别在于数据来源不同:
| API | 输入 Tensor 语义 | 典型场景 |
|---|---|---|
distribute_tensor() |
输入被当成 global tensor,rank 0 通常作为 source of truth | 初始化参数、分发输入、把普通 Tensor 转成 DTensor |
DTensor.from_local() |
输入就是当前 rank 的 local shard / replica / partial | 算子中间结果、已有分片权重、手动构造 local shard |
Method 1: distribute_tensor()
distribute_tensor() 适合从一个普通 torch.Tensor 创建 DTensor。它会根据 placements 把 global tensor 分发到 DeviceMesh 上:
|
|
如果 mesh 里有 4 个 rank,Shard(0) 会把 [8, 4] 的 global tensor 按第 0 维切成 4 份,每个 rank 的 local tensor 形状是 [2, 4]。
需要注意的是,distribute_tensor() 保持的是 single-device semantic:逻辑上它从一个完整 Tensor 出发,然后由 DTensor runtime 负责在 mesh 内 scatter / broadcast。实践中应把 rank 0 上的输入视为 source of truth,其他 rank 上的输入值不应该承载额外语义。
Method 2: DTensor.from_local()
DTensor.from_local() 适合每个 rank 已经有自己的 local tensor 的情况:
|
|
这里每个 rank 都提供一个 [2, 4] 的 local shard。如果 world_size = 4,那么 DTensor 的 global shape 会被解释为 [8, 4]。
run_check 的含义:
run_check=True:DTensor 会做一致性检查,例如 local tensor 的 shape / stride 是否能组成合法的 global tensor。更安全,但会引入额外通信。run_check=False:不做检查,直接相信调用方提供的 local tensor 和 placement 是正确的。更快,但调用方必须自己保证每个 rank 的 local shard 合法。
单脚本可运行示例
下面这个脚本不依赖 torchrun,直接用 torch.multiprocessing.spawn 启动多进程。机器有 GPU 时使用 cuda + nccl,否则退化到 cpu + gloo。
保存为 dtensor_create.py 后直接运行:
|
|
完整脚本:
|
|
输出里重点看三件事:
global_shape:用户看到的逻辑 Tensor shape,所有 rank 一致。placements:DTensor 当前的分布式布局。local_shape/local_value:当前 rank 实际持有的数据。
在 4 个 rank 上,distribute_tensor + Shard(0) 的输出大致是:
|
|
而 DTensor.from_local + Shard(0) 会保留每个 rank 预先构造的 local shard:
|
|
这正是两个 API 的核心差异:distribute_tensor() 负责从 global tensor 分发数据;DTensor.from_local() 只是把已有 local tensors 标注成同一个 global DTensor 的不同分片。
Working with DTensor
拿到 DTensor 之后,用户代码通常仍然按普通 torch.Tensor 的方式写:
- 直接调用 PyTorch 算子,例如
+、relu、matmul、sum。 - 用
to_local()查看当前 rank 的 local shard / replica。 - 用
full_tensor()收集完整 global tensor。 - 用
redistribute()显式改变 DTensor 的 layout。
DTensor 的关键价值在于:算子看到的是 global Tensor 语义,但实际执行时会尽量在 local tensor 上完成计算,并在必要时自动插入 collective communication。
Automatic Operator Parallelization
大多数 PyTorch 算子可以直接作用在 DTensor 上:
|
|
DTensor 会根据输入 placements 推导输出 placements:
- 对 elementwise 算子,如果输入都是相同的
Shard(0),输出通常仍然是Shard(0),每个 rank 只算自己的 local shard。 - 对 reduction 算子,如果 reduction 维度正好是 sharded 维度,输出可能变成
Partial("sum"),表示每个 rank 只有部分归约结果。 - 当后续算子需要完整值时,DTensor 会通过
all-reduce、all-gather等通信把 layout 转成需要的形式。
Local Tensor and Full Tensor
to_local() 返回当前 rank 实际持有的 local tensor:
|
|
如果 DTensor 是 Shard(0),那么 local_tensor 是当前 rank 的分片;如果 DTensor 是 Replicate(),那么 local_tensor 是当前 rank 的完整副本。
full_tensor() 会把所有 shards 收集起来,返回完整的 logical tensor:
|
|
这一步通常会触发通信。例如 Shard(0) -> Replicate() 本质上需要 all-gather,所以它适合调试、保存、验证,不适合放在训练 hot path 里频繁调用。
Redistributing DTensor
redistribute() 用来显式改变 DTensor 的 placements:
|
|
常见 layout 转换和通信关系:
| From | To | 可能触发的通信 |
|---|---|---|
Shard(dim) |
Replicate() |
all-gather |
Replicate() |
Shard(dim) |
local chunk / narrow |
Shard(src_dim) |
Shard(dst_dim) |
all-to-all 或组合通信 |
Partial() |
Replicate() |
all-reduce |
Partial() |
Shard(dim) |
reduce-scatter |
从使用习惯上看,redistribute() 是 DTensor 里非常重要的显式同步点:当自动 layout propagation 推导出的 layout 不是后续计算想要的形式时,就需要手动指定目标 placements。
单脚本测试代码
下面的脚本可以直接保存为 dtensor_working.py 运行:
|
|
它会测试四件事:
- DTensor elementwise 算子结果和普通 Tensor 一致。
sum(dim=0)这类跨 shard 维度的 reduction 可以得到正确 global 结果。to_local()和full_tensor()分别对应 local view 和 global view。redistribute()可以在Shard(0)、Replicate()、Shard(1)之间转换,并保持 global tensor 不变。
完整脚本:
|
|
预期输出里可以重点观察 placements 的变化:
|
|
这个例子里,sharded0 和 elementwise 都保持 Shard(0),说明 elementwise 计算可以直接在 local shard 上完成;reduced_replicated 变成 Replicate(),说明跨 shard 维度做 reduction 后需要归约到完整结果;sharded1 则展示了显式 reshard 的结果。
Automatic Parallelization in Action
前面几节介绍了 DTensor 的 layout 表达方式。本节关注另一个问题:当我们真的执行 PyTorch 算子时,DTensor 如何自动决定输出 placement,以及什么时候需要插入通信。
核心流程可以概括为:
- 用户调用普通 PyTorch 算子,例如
torch.sin(x)、torch.matmul(a, b)、x.sum()。 - DTensor 拦截算子调用,读取输入 DTensor 的
DTensorSpec。 - 根据算子的 sharding rule 推导合法的输出 placement。
- 如果当前输入 placement 不满足某个合法策略,就先对输入做
redistribute()。 - 在 local tensor 上执行真实算子,并把结果重新包装成 DTensor。
Example: Elementwise Operations
Elementwise 算子通常是最简单的情况。如果输入是 Shard(0),输出通常仍然是 Shard(0),每个 rank 只处理自己的 local shard:
|
|
逻辑上,用户看到的是完整 Tensor:
|
|
实际执行时,每张 GPU 只计算自己的 rows:
|
|
这类算子一般不需要通信,因为每个输出元素只依赖同位置的输入元素。
Example: Matrix Multiplication
矩阵乘法的 placement 决策取决于 shard 的维度是否落在输出维度或 contracting 维度上。
先看一个不需要通信的例子:
|
|
A @ B 的数学语义是:
|
|
这里 A 沿 M 维度做 Shard(0),而 B 在每个 rank 上都有完整副本。因此每个 rank 都能独立计算自己负责的输出 rows:
|
|
输出 placement 可以自然保持 Shard(0),不需要额外通信。
Example: Automatic Redistribution
再看一个需要自动重分布的例子:
|
|
初始状态:
|
|
Step 1: D = torch.matmul(A, B)
对 matmul(A, B) 来说:
A的第 0 维是输出 rows,Shard(0)是自然可保留的 output sharding。B的第 0 维是 contracting dimensionK,但当前B也被Shard(0)切开。- 每个 rank 只有一部分
K,无法直接算完整的A_local @ B。
因此,DTensor 需要先把 B 从 Shard(0) 转成 Replicate(),本质上触发一次 all-gather:
|
|
之后每个 rank 可以独立计算自己的输出 rows:
|
|
Step 2: E = torch.matmul(D, C)
此时:
D是Shard(0),沿输出 rows 切分。C是Replicate(),每个 rank 都有完整权重。
这又回到了理想情况,不需要通信:
|
|
输出 E 仍然是 Shard(0)。
Step 3: loss = E.sum()
E.sum() 是对整个 global Tensor 求和。由于 E 沿第 0 维切分,每个 rank 只能先算自己的局部和:
|
|
为了得到 global scalar,需要一次 all-reduce(sum):
|
|
因此,reduction 算子是最容易触发 Partial -> Replicate 或 Partial -> Shard 通信的场景。
Placement Decision Rules
下面是一些常见算子的直觉规则:
| Operation | Input placements | Output placement | 说明 |
|---|---|---|---|
| elementwise | inputs placements match | same as input | local shard 独立计算 |
matmul(Shard(0), Replicate()) |
rows sharded + full weight | Shard(0) |
无通信 |
matmul(Shard(0), Shard(0)) |
rows sharded + K sharded | usually Shard(0) |
需要先 gather 右矩阵 |
matmul(Shard(1), Shard(0)) |
K 维被两边切分 | Partial("sum") |
每个 rank 得到部分累加结果 |
sum() on Shard(dim) |
reduced sharded dimension | Partial("sum") then Replicate() |
需要 all-reduce 得到完整值 |
sum(dim=d) where d != shard_dim |
non-sharded dimension | Shard(adjusted_dim) |
通常不需要跨 rank 通信 |
这些规则不是硬编码在用户代码里的,而是由 DTensor 的 sharding propagation 机制根据每个 operator 的 schema 决定。
How DTensor Makes Decisions
DTensor 的自动并行不是全图优化,而是 eager execution 下的逐算子决策:
- 每个算子有一组合法的 sharding strategies,描述输入 placements 和输出 placements 的组合。
- 当当前输入 placements 已经匹配某个 strategy 时,算子可以直接执行。
- 当不匹配时,DTensor 会估算不同重分布方案的 communication cost,并选择一个局部成本较低的策略。
- 输入会先被
redistribute()到目标 layout,再执行 local operator。
相关源码入口可以从这些模块理解:
torch/distributed/tensor/_sharding_prop.py:负责 sharding propagation,推导 outputDTensorSpec。torch/distributed/tensor/_op_schema.py:定义OpSchema、OpStrategy等策略表达。torch/distributed/tensor/_ops/:不同 operator 的 sharding rule。torch/distributed/tensor/_dispatch.py:DTensor 的 dispatch 流程,在需要时调用 redistribution。
这也带来一个限制:DTensor 当前更接近「per-operator greedy decision」,不是像 XLA GSPMD 那样对整张计算图做全局最优 sharding 规划。因此,局部通信最少的选择不一定是整个模型端到端通信最少的选择。实际训练中仍然需要用户通过合适的初始 placement、redistribute()、以及更高层的 TP/FSDP 策略来引导 layout。
Comparison with JAX GSPMD
JAX 的 GSPMD / XLA 通常会在 compile-time trace 完整计算图,然后为全图 sharding 求解更全局的方案。PyTorch DTensor 则更贴近 PyTorch eager 模型,运行时逐算子做 placement propagation 和 redistribution。
| Aspect | PyTorch DTensor | JAX GSPMD / XLA |
|---|---|---|
| Optimization scope | Per-operator, local greedy decision | Whole-graph optimization |
| Decision time | Runtime eager execution | Compile-time tracing |
| Cost model | Lightweight redistribution cost | Global cost model / solver |
| User control | placements + redistribute() |
with_sharding_constraint / partition spec |
| Strength | 更贴近 PyTorch eager,调试和增量迁移更直接 | 全图视角更容易得到全局更优 sharding |
| Trade-off | 可能出现局部最优但全局非最优的通信 | 编译成本和约束表达更重 |
总结来说,DTensor 的自动并行能让很多算子在不改用户代码的情况下自动分布式执行,但它不是完全替代并行策略设计。更好的理解方式是:DTensor 提供了统一的 layout 表达、算子级 propagation 和必要通信插入;用户仍然需要在模型并行、数据并行、参数布局这些更高层做结构化设计。
Guiding Placement Decisions with Explicit Redistribution
用户不能直接覆盖某个 PyTorch operator 的 DTensor sharding rule,也不能在调用 torch.matmul() 时手动指定“这个 op 的输出一定要是什么 placement”。但实际使用中有一个很重要的控制手段:在 operator 之前插入显式的 redistribute()。
这相当于给 DTensor 一个 placement hint:虽然你不能改 op schema 的选择逻辑,但你可以控制这个 operator 实际看到的输入 layout。由于 DTensor 是逐算子做 placement propagation,这个技巧常常能避免后续不必要的 reshard。
The Problem: Greedy Per-Operator Decisions
考虑下面这个计算图:
|
|
对 matmul(A, B) 来说:
A是Shard(1),沿 columns 切分,也就是沿 contracting dimensionK切分。B是Shard(0),沿 rows 切分,同样也是沿 contracting dimensionK切分。- 每个 rank 可以计算一部分
K对输出的贡献,因此输出C的自然 layout 是Partial("sum")。
可以把这一步理解成:
|
|
接下来执行 torch.relu(C) 时,问题出现了:非线性算子不能直接作用在 Partial 上,因为:
|
|
所以 relu() 之前必须把 C 从 Partial("sum") resolve 成一个 concrete layout。由于 DTensor 在 eager 模式下只看到当前这个 relu(C) operator,它并不知道后面马上会有一个 D + F,也不知道 F 已经是 Shard(1)。
如果 DTensor 在 relu() 前把 C resolve 成 Replicate(),那么通信路径可能变成:
|
|
这就是 per-operator greedy decision 的局限:单个 operator 看起来合理的 layout,不一定是整个后续计算图通信最少的 layout。
The Solution: Use redistribute() as a Placement Hint
用户通常知道更长的计算上下文。因此,可以在 relu() 前显式把 C resolve 成后续更需要的 layout:
|
|
这里 C.redistribute(mesh, [Shard(1)]) 的语义是:把每个 rank 上的 partial output 沿第 1 维做 reduce-scatter,直接得到按 columns 切分的完整结果:
|
|
这样后续两步都可以保持 Shard(1):
torch.relu(Shard(1)) -> Shard(1):elementwise 算子在 local shard 上执行,不需要通信。Shard(1) + Shard(1) -> Shard(1):两个输入 layout 匹配,也不需要额外 reshard。
Communication Trade-Off
显式 redistribute() 不是“免费优化”,它只是让用户选择通信发生在哪里、以及通信后的 layout 长什么样。
在上面的例子里,有两种可能路径:
| Strategy | Communication | Downstream layout |
|---|---|---|
Let DTensor resolve Partial -> Replicate before relu |
all-reduce over full [1024, 256] |
后续和 F: Shard(1) 计算时可能还要 reshard |
User explicitly resolves Partial -> Shard(1) |
reduce-scatter over [1024, 256] |
后续 relu 和 + F 都保持 Shard(1) |
选择哪一种取决于后续计算和 tensor size。如果后面大量算子都希望 Replicate(),那么 all-reduce 可能是合理的;如果后面会和 Shard(1) 的张量继续计算,那么提前 reduce-scatter 到 Shard(1) 往往更好。
因此,redistribute() 的价值不只是“改变 layout”,而是把用户对后续计算图的全局理解显式写进 eager 程序里。
Comparison with JAX with_sharding_constraint()
JAX 提供了更 declarative 的方式:jax.lax.with_sharding_constraint() 可以告诉 XLA 某个中间结果期望采用什么 sharding。编译器在看到完整图之后,可以把这个 constraint 纳入全局优化。
|
|
两者的差异可以总结为:
| Aspect | PyTorch DTensor redistribute() |
JAX with_sharding_constraint() |
|---|---|---|
| Style | Imperative layout transformation | Declarative compiler constraint |
| Timing | 立即执行,立即通信 | 编译期参与全图优化 |
| Scope | 影响下一个及后续 eager operators 看到的 input placement | 影响 XLA 对整张图的 sharding 规划 |
| User control | 直接、显式、容易调试 | 更抽象,但可能获得更全局的优化 |
所以,在 PyTorch DTensor 里,“引导 placement 决策”的核心方式就是:在关键边界上显式插入 redistribute(),把中间结果转成后续计算最自然的 layout。
实现原理 __torch_dispatch__
PyTorch 算子下发流程:
-
OneFlow: Redesign the Distributed Deep Learning Framework from Scratch, https://arxiv.org/pdf/2110.15032 ↩︎
-
GSPMD: General and Scalable Parallelization for ML Computation Graphs, https://arxiv.org/pdf/2105.04663 ↩︎
Author Houmin Wei
Publish January 1, 0001
LastMod August 27, 2026
License 本作品采用 CC BY-NC-ND 4.0 许可协议进行许可,转载时请注明原文链接
如果你在浏览博客的过程中发现了任何问题,欢迎在对应文章下评论。如果你有其他事情想要咨询,可以通过邮件联系我。