DeviceMesh 与张量并行
本教程共 60 篇 · 第 58 篇 · 更新于 2026-08-17 · 约 4 分钟阅读
本节目标:搞清并行策略怎么选,学会用 DeviceMesh 组织多卡设备,理解张量并行与流水线并行的核心思想。
并行策略全家桶
前面两章已经见了两位选手:DDP 每卡一份完整模型,FSDP 把参数分片。它们都是数据并行——所有卡算同一份模型的不同数据。
当模型大到数据并行也扛不住时,得上模型并行(model parallelism):把模型本身拆开,不同部分放不同卡。主流拆法有两种:
- 张量并行(Tensor Parallel,TP):把一层的权重矩阵切开,几张卡各算一部分,再拼结果。像把一个大矩阵乘法分给几个人同时算。
- 流水线并行(Pipeline Parallel,PP):按层把模型切成几段,卡 0 算前几层,卡 1 算中间几层。像工厂流水线,上一段的输出是下一段的输入。
两种并行的差别可以这么记:数据并行是每个老师手里同一本书,学生分工批改不同作业;张量并行是一本书拆成几本分册,每人拿一册一起念;流水线并行则是书按章节分给不同人,念完一章传下一章。
什么时候用哪个?官方给了一条很实用的选择链:
- 模型放得进单卡,想多用几张卡加速 → DDP(§56);
- 模型放不进单卡 → FSDP2(§57);
- FSDP 也到极限了 → 加张量并行;
- 单卡连一段模型都塞不下 → 再上流水线并行。
Note为什么 FSDP 到几百张卡会吃力?all-gather 这类集合通信有环形延迟,卡越多延迟越明显。张量并行把 FSDP 的通信范围缩小到主机之间,延迟成本能降一个量级。
DeviceMesh:设备的棋盘
组合多种并行策略,第一件事是搞清设备怎么分组。比如 16 张卡,张量并行要 8 张卡一组,数据并行要跨主机配对。手写 dist.new_group 得自己算每个 rank 属于哪个组,又绕又容易错。
DeviceMesh(设备网格)就是干这个的:把设备组织成多维网格,每个维度对应一种并行策略,底层进程组自动建好。
from torch.distributed.device_mesh import init_device_mesh
# 一维网格:8 张卡,一个组
tp_mesh = init_device_mesh("cuda", (8,))
# 二维网格:2×4,两个维度各有一个进程组
mesh_2d = init_device_mesh("cuda", (2, 4), mesh_dim_names=("replicate", "shard"))
# 按名字取底层的进程组
shard_group = mesh_2d.get_group(mesh_dim="shard")
三维网格也支持,还能从父网格切出子网格,复用已建好的通信器:
mesh_3d = init_device_mesh("cuda", (2, 2, 2), mesh_dim_names=("replicate", "shard", "tp"))
hsdp_mesh = mesh_3d["replicate", "shard"] # 切出二维子网格
tp_mesh = mesh_3d["tp"] # 切出一维子网格
有了 DeviceMesh,混合分片数据并行(HSDP,主机内 FSDP + 主机间 DDP)两行搞定:
import torch.nn as nn
from torch.distributed.fsdp import fully_shard
class ToyModel(nn.Module):
def __init__(self):
super().__init__()
self.net1 = nn.Linear(10, 10)
self.relu = nn.ReLU()
self.net2 = nn.Linear(10, 5)
def forward(self, x):
return self.net2(self.relu(self.net1(x)))
mesh_2d = init_device_mesh("cuda", (2, 4), mesh_dim_names=("dp_replicate", "dp_shard"))
model = fully_shard(ToyModel(), mesh=mesh_2d)
张量并行:把矩阵切开算
张量并行最早来自 Megatron-LM 论文,核心是切矩阵乘法。一个线性层 y = xW,可以把 W 按列切成两块,两张卡各算一半输出再拼起来;也可以按行切,两张卡各算一半输入,最后结果相加。
PyTorch 提供了现成的并行风格(ParallelStyle),照着层名配置就行:
ColwiseParallel:按列切,适合 MLP 的第一层、注意力的 q/k/v 投影;RowwiseParallel:按行切,适合 MLP 的输出层、注意力的输出投影;SequenceParallel:序列并行,归一化层按序列维度分片,省激活内存。
以 Transformer 的 FeedForward 层为例:w1、w3 列切,w2 行切,三层算完只通信一次:
from torch.distributed.tensor.parallel import (
ColwiseParallel,
RowwiseParallel,
parallelize_module,
)
tp_mesh = init_device_mesh("cuda", (8,)) # 主机内 8 卡
layer_tp_plan = {
"feed_forward.w1": ColwiseParallel(),
"feed_forward.w2": RowwiseParallel(),
"feed_forward.w3": ColwiseParallel(),
}
for layer in model.layers:
parallelize_module(module=layer, device_mesh=tp_mesh, parallelize_plan=layer_tp_plan)
通信是自动的。parallelize_module 把参数换成 DTensor,需要时自动插入 all-reduce 等通信操作,你只管声明每层怎么切。注意力层的 q/k/v 投影同理,列切后注意多头维度的对应关系,用 use_local_output=False 保持 DTensor 输出即可。
序列并行是张量并行的常见搭档:把 LayerNorm、Dropout 这类层的计算按序列维度分片,激活不用在每个 Transformer 块之间来回收集,能省一大块显存。大规模训练里激活内存常常先爆,所以 TP 训练通常搭配序列并行(SequenceParallel)。
输出层还可以做损失并行(loss parallel):模型输出按词表维度分片时,用 loss_parallel() 上下文包住交叉熵计算,不用把整个输出收集到每张卡,省显存又省通信。
Tip张量并行通常只在主机内做,因为每层都要通信,走 NVLink 才快。跨主机的部分交给 FSDP 或 DDP。层内通信频繁是它的特点,也是它的局限。
张量并行 + FSDP:2D 并行
大模型训练的标准姿势:主机内张量并行,主机间 FSDP。一个二维网格就能表达,各取所需:
mesh_2d = init_device_mesh("cuda", (8, 8)) # 8 台机器 × 8 卡
tp_mesh = mesh_2d["tp"] # 主机内 8 卡一组,做张量并行
dp_mesh = mesh_2d["dp"] # 跨主机维度,做 FSDP
model_tp = parallelize_module(model, tp_mesh, tp_plan)
model_2d = fully_shard(model_tp, mesh=dp_mesh)
这是 TorchTitan 等框架 3D 并行的基础。再加一层流水线并行,就是完整的 3D 并行:流水线切层间,张量并行切层内,FSDP 切参数。三个维度各管一摊,互不干扰。
流水线并行:接力跑
流水线并行把模型按层切成几段,每段放一张卡。前向像接力:卡 0 算完传给卡 1。但直接跑有个浪费——某一时刻只有一张卡在算,其他卡干等。
解决办法是微批量(micro-batch):把一个大 batch 切成多个小块,让不同的小块错开在不同阶段上跑,像流水线一样把卡填满。气泡(bubble)是流水线固有的空转,GPipe、1F1B 这些调度策略就是尽量缩小它。
PyTorch 的流水线 API 在 torch.distributed.pipelining,由 PiPPy 项目并入核心:
from torch.distributed.pipelining import pipeline
pipe = pipeline(
module=model,
num_chunks=4, # 微批量数量
example_args=(example_input,),
split_points=["layers.4", "layers.8"], # 切两刀,分 3 段
)
Warning
torch.distributed.pipelining尚在测试阶段,API 可能变化。入门阶段理解概念即可,真要大规模训练,先用 TorchTitan 这类封装好的框架。
小结
这一章的三张牌:DeviceMesh 管设备组织,张量并行拆层内计算,流水线并行拆层间顺序。选型记住那条链:单卡装得下用 DDP,装不下用 FSDP,还不够加 TP,层太多上 PP。下一章把前面所有技巧串成方法论——性能调优指南。