第 6 章 分布式训练基础
第6章 分布式训练基础
本章界定分布式训练的基本概念与通信模型,梳理数据并行(DDP、FSDP)、张量并行、序列并行与自动并行化的演进路 径与核心原理,并以 FSDP 生产配置实践收尾。推理侧的并行优化不在本章展开。
6.1 分布式训练概念框架
分布式训练的核心目标是将单机无法容纳的模型或数据规模扩展到多机多卡上高效执行。理解分布式训练,需要先建立统 一的概念框架。
6.1.1 并行策略分类
当前主流的分布式训练可以按四个维度分类,如图6-1所示。 Distributed Training Strate gy Data Parallel Model Parallel Pipeline Parallel Hybrid Parallel Data Parallelism Model Parallelism Pipeline Parallelism Hybrid Parallelism DDP / FSDP / ZeRO Tensor Parallel Expert Parallel GPipe / PipeDream / 1F1B 3D / 4D / 5D Parallelism Tensor Parallelism Expert Parallelism 图6-1 分布式训练策略分类体系
6.1.2 通信模型对比
分布式训练中,Worker 之间需要同步梯度或参数。历史上出现过两种主流通信模型: •参数服务器架构:Worker 负责计算梯度,Server 负责存储和更新参数,Worker 将梯度 Push 到 Server,再从 Server Pull 最新参数。优势是支持异步更新、对落后者容忍度高,代价是通信瓶颈——所有梯度和参数都经过 PS 节点。 •AllReduce 集合通信架构:所有 Worker 对等地参与通信,通过 AllReduce 操作直接聚合所有 Worker 的梯度。现代训 练框架(PyTorch DDP、Horovod)普遍采用此方案。
6.1.3 同步与异步训练
•同步训练:所有 Worker 在每一步都等待彼此完成梯度计算,对梯度做全局 AllReduce 后统一更新参数。保证严格的数 学等价性,但存在“落后者问题”,最慢的 Worker 决定整体速度。 •异步训练:Worker 独立计算梯度并更新参数,无需等待。吞吐高但引入梯度陈旧(Gradient Staleness)问题——某个 Worker 读取的参数可能已被其他 Worker 更新过多次。Mitliagkas 等人(2016)证明梯度陈旧等价于在梯度上施加额 外动量项,可能影响收敛质量。
6.1.4 通信复杂度与 AllReduce
分析通信开销有两个经典模型,加上 Ring 与树式算法的取舍,共同构成评估集合通信的基础。
- α-β 模型与LogP模型 •α-β 模型:通信时间 T = α + β ⋅ M ,其中 α 是延迟(latency),β 是带宽的倒数(1/B),M 是消息大小。该模型假设 网络无拥塞、带宽恒定。 •LogP 模型:扩展为 L(延迟)、o(overhead,CPU 处理开销)、g(gap,连续发送间隔)、P (处理器数),更精确地 建模实际网络行为。
- Ring AllReduce算法 Ring AllReduce 将 P 个节点组织成逻辑环,通过 Reduce-Scatter 与 All-Gather 两阶段完成梯度聚合。算法带宽最优,总 通信量接近 2M ,不随节点数增长;延迟则随环上步数线性增加,适合大消息的梯度同步。 以 P 个 GPU 聚合大小为 M 的梯度为例,Ring 算法的耗时可以拆成两项:一是带宽项,每个节点沿环实际收发的总量约 为 2M ,几乎不随 P 变化,除以链路带宽即为主体耗时;二是延迟项,随环上步数(正比于 P )线性增长。因此在大消息 场景(如全模型梯度同步)下,Ring AllReduce 接近带宽最优,规模扩大主要侵蚀延迟项而非带宽项。
- 带宽与延迟权衡 AllReduce 算法选择的本质是带宽与延迟的权衡:环式算法以线性增长的步数为代价,换取接近理论下限的总流量;树式 算法则用更高的总流量换取对数级增长的延迟步数。NCCL 会根据消息大小和 GPU 数在多种算法间自动选择,应用层通 常无需干预。
6.1.5 分布式初始化
最基础的 PyTorch 分布式初始化代码示例如下:
import os
import torch
import torch.distributed as dist
def init_distributed():
# Method 1: environment variables (MASTER_ADDR, MASTER_PORT, WORLD_SIZE, RANK)
dist.init_process_group(
backend='nccl', # NCCL preferred for GPU communication
init_method='env://')
# Method 2: TCP initialization
# dist.init_process_group(backend='nccl', init_method='tcp://master_ip:23456',
# rank=rank, world_size=world_size)
local_rank = int(os.environ.get('LOCAL_RANK', 0))
torch.cuda.set_device(local_rank)
return dist.get_rank(), dist.get_world_size()•backend 选择: nccl (NVIDIA GPU 集合通信库)是 GPU 场景的最佳选择; gloo (Facebook)可用于 CPU 或跨平 台场景; mpi 适合 HPC 环境。
6.1.6 关键通信原语
PyTorch 分布式模块封装了 NCCL 的集合通信操作,各原语的功能与复杂度如表6-1所示。 表6-1 关键通信原语 原语 功能 复杂度 all_reduce 全局求和/平均后广播 O(M ) 带宽,O(P ) 或 O(log P ) 延迟 all_gather 收集所有rank的数据 O(PM ) 接收量 reduce_scatter 规约后分散到各rank O(M ) 接收量 broadcast 从root广播 O(M ) 接收量,O(log P ) 延迟 all_to_all 每个rank分发送到全部rank O(PM ),MoE场景关键操作
6.1.7 通信计算比分析
- 通信-计算比定义 定义通信-计算比(Communication-to-Computation Ratio): Tcomm R=
Tcomp 当 R > 1 时说明通信是瓶颈。理想情况下 R < 0.1 才能实现良好的弱扩展。大模型训练中,通信优化(梯度压缩、通信- 计算重叠)是持续研究的核心课题。 2) 硬件量化估算 抽象的 R < 0.1 需要落到具体硬件才有操作意义。H100 集群的典型带宽基线:节点内 GPU-GPU(NVLink 4/NVSwitch) 单向约 450 GB/s;节点间 InfiniBand NDR 400 有效约 40 GB/s(含 NCCL 开销)。 •示例 A:7B 模型 DDP,8 卡单节点:FP16 梯度 ≈ 14 GB,Ring AllReduce 每卡通信量 ≈ 24.5 GB,NVLink 450 GB/s 下 T_comm ≈ 54ms;H100 一步计算约 200-400ms(bs=32,seq=2048),R ≈ 0.14–0.27,已在瓶颈边界附近。 •示例 B:70B 模型 DDP,64 卡跨 8 节点:梯度 ≈ 140 GB,Ring AllReduce 每卡约 277 GB;IB 有效带宽 40 GB/s, T_comm ≈ 6900ms,远超计算时间——纯 DDP 在此规模完全不可行,必须结合 ZeRO/FSDP(每卡同步梯度分片 ≈ 2.2 GB)或改用 Tensor/Pipeline 并行。这一估算揭示了规模临界点:单节点内(NVLink 域)DDP 可行到约 30B 参 数,跨节点(IB 域)超过 10B 参数就需要切换通信策略。 参考文献:Li et al., “Scaling Distributed Machine Learning with the Parameter Server”, OSDI 2014。
6.2 数据并行到 AllReduce DDP
数据并行(Data Parallelism, DP)是最基本也最广泛使用的分布式训练策略:每个 Worker 持有模型的完整副本,处理 不同的数据子集,然后同步梯度。
6.2.1 参数服务器演进
参数服务器(Parameter Server)由 Li 等人在 OSDI 2014 论文中系统化提出,其核心设计包括: •分布式键值存储抽象:参数按 Key 分片存储在多个 Server 节点上,Worker 通过 Push/Pull 接口读写。 •一致性模型:支持多种一致性协议:BSP(Bulk Synchronous Parallel)、SSP(Stale Synchronous Parallel)、ASP (Asynchronous Parallel)。 •容错:Server 端基于一致性哈希的复制,Worker 端基于快照恢复。 •通信优化:Range Push/Pull 将小消息合并发送,利用消息合并降低延迟开销。 PS 的致命弱点在于通信拓扑不对称:W 个 Worker 与 S 个 Server 之间形成 W × S 的全连接通信。当 W 很大时, Server 端的网卡(NIC)成为瓶颈——一个 100Gbps 的 NIC 同时服务 64 个 Worker 时,每 Worker 带宽仅约 1.5Gbps。
6.2.2 梯度聚合策略
从 PS 到 AllReduce,梯度聚合经历了多种策略,如表6-2所示。 表6-2 梯度聚合策略对比 策略 通信模式 优点 缺点 PS Push-Pull 多对多 灵活的一致性模型 Server 瓶颈 Tree AllReduce 树形递归聚合 延迟 O(log P ) 根节点带宽瓶颈 Recursive Halving Doubling 递归折半与翻倍 低延迟 实现复杂 Ring AllReduce 环形流水线 带宽最优 延迟 O(P ) Hierarchical AllReduce 两级聚合 匹配物理拓扑 需要拓扑感知
6.2.3 DDP 核心机制
DDP 是 PyTorch 官方推荐的数据并行解决方案,自 v1.5 起取代了旧的 DataParallel (单机多卡,存在 Python 线程竞 争和 GIL 问题)。DDP 于 v1.0 引入,在 v1.5 完成重大架构重写(Reducer、梯度桶化、通信重叠),成为官方推荐方案。 DDP 的核心机制: •多进程架构:每个 GPU 运行独立进程,完全消除 Python GIL 限制。 •梯度同步内置:通过 autograd hook 在每次 backward() 后自动触发 AllReduce。 •Reducer 模块:负责梯度桶化(bucketing)和异步通信。
6.2.4 Reducer 与梯度桶化
- Reducer 梯度同步流程 DDP 的核心是 Reducer 类( torch/csrc/distributed/c10d/reducer.cpp ): // Simplified Reducer workflow void Reducer::autograd_hook(int index) { // 1. mark this gradient ready mark_variable_ready(index);
// 2. when all gradients in a bucket are ready
if (bucket.all_ready) {
// 3. launch async AllReduceallreduce_bucket(allreduce_bucket);
}
// 4. at the end of backward, wait for all outstanding AllReduce
// (usually earlier ones complete during last bucket communication)
}- 梯度桶化 DDP 将参数的梯度分组到多个桶(bucket)中。默认桶大小为 25MB。桶化策略: •反向桶重建:DDP 根据 backward() 的执行顺序动态分配参数到桶中。先产生梯度的参数放在前面的桶,便于尽早启 动通信。 •桶大小权衡:较大的桶减少 AllReduce 次数(降低延迟开销),但推迟了通信启动时机;较小的桶反之。 model = torch.nn.parallel.DistributedDataParallel( model,
bucket_cap_mb=50, # bucket capacity 50MB
gradient_as_bucket_view=True, # avoid gradient copy, use view
static_graph=True, # enable when computation graph is static, reduce overhead)
6.2.5 通信计算重叠
- Backward 与 AllReduce 重叠 DDP 最关键的优化是实现 backward() 与梯度 AllReduce 的通信重叠,如图6-2所示。 Backward Communication layer L backward (grad_L ready) AllReduce(grad_L) starts layer L-1 backward (overlapping with comm) AllReduce(grad_L) completes AllReduce(grad_L-1) starts layer L-2 backward Backward Communication 图6-2 DDP 中 Backward 与 AllReduce 的流水线重叠 这种重叠减少的等待时间约等于 T −T 中未被重叠部分。实际效果取决于计算量和通信量的比率。 backward allreduce
- barrier 的代价 全局屏障(barrier)强制所有进程同步到同一点。滥用 barrier 是性能杀手:
# Anti-pattern: barrier on every step
for step in range(total_steps):
loss.backward()
dist.barrier() # Danger! straggler blocks everyone
dist.all_reduce(gradients) # all_reduce itself implies synchronization
optimizer.step()NCCL 中 barrier 的实现是连续的环状数据传输,时间约为 O(P ⋅ α)。在 1000+ GPU 规模下,每次 barrier 可能耗时上百 毫秒。
6.2.6 梯度累积实践
- 梯度累积与 no_sync 当显存不足以支持所需的有效 batch size 时,梯度累积是标准解决方案: for micro_step, micro_batch in enumerate(micro_batches): with (model.no_sync() if micro_step < accumulation_steps - 1 else contextlib.nullcontext()):
# no_sync() disables gradient sync (trigger AllReduce only at last step)
loss = model(micro_batch) / accumulation_steps
loss.backward()
# AllReduce triggers after the last micro-step
optimizer.step()
optimizer.zero_grad()model.no_sync() 是 PyTorch DDP v1.1+ 引入的 context manager,它临时禁用梯度同步,避免了每个 micro-batch 后都执行不必要的 AllReduce。等效 global batch size = micro_batch_size × accumulation_steps × world_size。 2) Horovod 设计哲学 Horovod 由 Uber 开发,将分布式训练抽象为统一接口。其核心创新在于 hvd.allreduce() 调用模式和 Tensor Fusion:将多个小张量自动拼接成大块再执行 AllReduce,最大化带宽利用率。
import horovod.torch as hvd
hvd.init()
torch.cuda.set_device(hvd.local_rank())
optimizer = hvd.DistributedOptimizer(optimizer, named_parameters=model.named_parameters() ) hvd.broadcast_parameters(model.state_dict(), root_rank=0) 但随着 PyTorch DDP 的日趋成熟(内置 Reducer、通信重叠、no_sync),DDP 已成为数据并行的首选方案。
6.2.7 内存天花板
DDP 每张卡持有完整副本,在模型规模扩大后很快触及单卡显存上限。以 7B 参数模型(bf16 权重 + Adam 优化器)为 例,显存占用如表6-3所示。 表6-3 7B 模型 DDP 训练显存占用 项目 数据类型 显存占用 模型参数 bf16 14 GB 梯度缓冲 bf16 14 GB 主权重(AMP) fp32 28 GB Adam momentum fp32 28 GB Adam variance fp32 28 GB 合计(无 activation) — 112 GB H100 SXM5 单卡 80 GB,纯 DDP 无法放下 7B 模型的完整训练状态。决策边界: •参数 × 16 ≤ 单卡显存:使用 DDP,零额外通信开销。混合精度 Adam 训练下每参数占 16 字节(BF16 权重与梯度共 4 字节,FP32 主权重与 Adam 动量/方差共 12 字节)。 •参数 × 16 > 单卡显存但参数本身放得下:用 ZeRO-1(PyTorch ZeroRedundancyOptimizer ),保留 DDP 梯度同步 语义,仅将优化器状态分片(显存除以 GPU 数)。8 卡 H100 训练 7B:优化器状态从 56 GB 降至 7 GB/卡,合计 63 GB,勉强可行。 •参数本身超出单卡:必须使用 FSDP / ZeRO-3 这类全分片方案。
# ZeRO-1 via PyTorch ZeroRedundancyOptimizer
from torch.distributed.optim import ZeroRedundancyOptimizer
optimizer = ZeroRedundancyOptimizer(
model.parameters(),
optimizer_class=torch.optim.AdamW,
lr=1e-4,)
Model weights and gradients remain local (not sharded);
only optimizer states are sharded across ranks
参考文献:Li et al., “PyTorch Distributed: Experiences on Accelerating Data Parallel Training”, VLDB 2020。 Sergeev & Del Balso, “Horovod”, arXiv 2018。
6.3 全分片数据并行 FSDP
标准数据并行(DDP)的核心问题是:每个 Worker 持有完整的参数、梯度和优化器状态副本。以 Adam 优化器训练的 Llama-70B 为例: •参数(BF16):70B × 2 bytes = 140 GB •梯度(BF16):140 GB •Adam 优化器状态(FP32 master weights + momentum + variance):70B × 4 × 3 = 840 GB •总计:单卡需约 1120 GB 显存,远超 A100(80GB)或 H100(80GB)的单卡容量。 即使用上数据并行,每张卡依然需要 840GB。FSDP(Fully Sharded Data Parallel)和 ZeRO-3 的核心思想是:将模型参 数、梯度和优化器状态在所有数据并行 Worker 之间分片存储。
6.3.1 ZeRO 分片定位
FSDP(Fully Sharded Data Parallel)是 ZeRO(Zero Redundancy Optimizer)思路在 PyTorch 中的实现。ZeRO 的核 心思想是在数据并行基础上,将优化器状态、梯度与参数分别在各卡间分片,消除每卡上的冗余副本——分片越彻底显存 越省,但通信开销越大。FSDP 采用 ZeRO-3 全分片:每卡仅保存 1/N 参数,前向/反向传播中需频繁 AllGather 参数,通 信量反而比 DDP 更大,以通信换显存。
6.3.2 分片状态与流程
FSDP 维护参数的三种状态,如图6-3所示: •全收集状态(All-Gathered):在正向/反向传播期间,需要完整参数进行计算, all_gather 将所有分片收集起来。 •分片存储状态(Sharded):计算完成后立即释放全量参数,仅保留自己的分片,释放的显存实现空间换取时间。 •重分片状态(Re-Sharded):梯度经过 reduce_scatter 后,每个 Worker 持有归约后的梯度分片。 GPU 0 GPU 1 GPU 2 GPU 3 Forward (Layer N) All-Gather W_N (shard▶full) Compute Fwd(W_N) Discard full W_N (keep shard) All-Gather W_N+1 (prefetch) Backward (Layer N+1 ▶ Layer N) All-Gather W_N+1 (reuse for bw) Compute Bwd(W_N+1) Reduce-Scatter grad_W_N+1 All-Gather W_N Compute Bwd(W_N) Reduce-Scatter grad_W_N GPU 0 GPU 1 GPU 2 GPU 3 图6-3 FSDP 正反向传播中参数收集、释放与梯度归约的时序
6.3.3 FSDP 配置参数
FSDP 的包装策略(Wrapping Policy)决定哪些子模块被包装为 FSDP 单元: from torch.distributed.fsdp import FullyShardedDataParallel as FSDP from torch.distributed.fsdp.wrap import ( size_based_auto_wrap_policy, transformer_auto_wrap_policy, )
Strategy 1: auto-wrap by parameter count
auto_wrap_policy = functools.partial( size_based_auto_wrap_policy, min_num_params=1e8 # 100M parameters )
Strategy 2: auto-wrap by transformer layer class
auto_wrap_policy = functools.partial( transformer_auto_wrap_policy, transformer_layer_cls={LlamaDecoderLayer, GPT2Block} ) model = FSDP( model,
auto_wrap_policy=auto_wrap_policy,
sharding_strategy=ShardingStrategy.FULL_SHARD,
cpu_offload=CPUOffload(offload_params=False),
mixed_precision=MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,),
backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
forward_prefetch=True,
limit_all_gathers=True,) 关键参数解析: • sharding_strategy : FULL_SHARD (ZeRO-3)、 SHARD_GRAD_OP (ZeRO-2)、 NO_SHARD (DDP)。 • backward_prefetch : BACKWARD_PRE 在当前层 backward 期间预取下一层的 All-Gather 参数,实现通信-计算重 叠。 • forward_prefetch :在 forward 期间预取下一层的 All-Gather 参数。 • limit_all_gathers :限制同时执行的 All-Gather 数量,避免占用过多 CUDA 临时内存。 • use_orig_params :使用原始参数名而非 FSDP 扁平化后的名称,方便 checkpoint 保存。
6.3.4 混合精度与 FSDP
FSDP 的混合精度设置比 DDP 更精细: MixedPrecision(
param_dtype=torch.bfloat16, # Parameter precision used in forward/backward
buffer_dtype=torch.bfloat16, # Buffer precision
reduce_dtype=torch.float32, # Precision for gradient Reduce-Scatter
keep_low_precision_grads=False, # Whether to keep low precision gradients)
6.3.5 通信模式分析
对一层参数 W ∈ R ,参数量 N = d ⋅ d 。FSDP 的每层通信由全参数的 All-Gather 与梯度的 Reduce-Scatter 构 din ×dout in out 成: •Forward:All-Gather 全量参数,每卡 N •Backward:All-Gather 全量参数(复用前向)+ Reduce-Scatter 梯度,每卡约 2N •每步总通信量约 3N /卡,与 DDP 的 2N 相比增加 50%,换来的收益是模型状态显存降至 1/P 配合 backward_prefetch ,下一层的参数 All-Gather 可与当前层的反向计算重叠,理想情况下可隐藏约 30-50% 的通信 延迟。 CPU Offload 参数卸载 FSDP 支持将参数或优化器状态卸载到 CPU: CPUOffload(offload_params=True) # Store parameters on CPU, DMA to GPU when needed 但这引入了 PCIe 带宽瓶颈(A100 PCIe Gen4 x16 仅约 64 GB/s 双向,即单向 32 GB/s),可能显著拖慢训练。通常仅在 极端内存受限场景使用。
6.3.6 混合分片并行
HSDP 是多节点场景的关键优化:节点内使用完整复制(减少通信),节点间使用分片: Node 1 (8 GPUs): Shard across 8 GPUs ▶ sub-group All-Gather Node 2 (8 GPUs): Shard across 8 GPUs ▶ sub-group All-Gather Then: Reduce-Scatter across 16 GPUs (two nodes) 在 PyTorch FSDP v2 中通过 HybridShardingStrategy 实现,需要指定 process_group 和 intra_node_pg 。
6.3.7 FSDP 与 ZeRO 的关系
FSDP 是 ZeRO-3 思路在 PyTorch 中的原生实现:通过 auto_wrap_policy 自动决定分片单元,与 PyTorch 生态无缝集 成;DeepSpeed ZeRO-3 以 JSON 配置驱动,CPU/NVMe 卸载能力更完整。两者在 ZeRO-3 级别语义等价,差异体现在 框架生态与卸载能力。
6.3.8 FSDPv2 新架构
PyTorch 2.2 引入了基于 DTensor 的新版 FSDP(FSDPv2, torch.distributed._composable.fsdp ),采用 SPMD 编 程范式:
from torch.distributed._composable.fsdp import fully_shard
from torch.distributed.tensor import Shard
# Mark parameters and gradients as sharded on dp mesh
for layer in model.layers:
fully_shard(layer, mesh=dp_mesh)
fully_shard(model, mesh=dp_mesh)FSDPv2 更模块化,与 DTensor 的 sharding 策略(Shard/Replicate)统一,也更易于与 TP/PP 组合。FSDP1 与 FSDPv2 的关键 API 差异如表6-4所示。 表6-4 FSDP1 与 FSDPv2 API 对比 维度 FSDP1 (PyTorch 2.0-2.4) FSDPv2 (PyTorch 2.2+) 入口 FullyShardedDataParallel wrapper fully_shard() + DeviceMesh 分片定义 ShardingStrategy enum DTensor Shard() placement 设备布局 隐式(单 PG) 显式 DeviceMesh,多维网格 混合并行 手动组合 TP/PP DTensor 原生跨 mesh 维度 计算图表示 eager 模式 支持 torch.compile 图捕获 API 稳定性 稳定( torch.distributed.fsdp ) 实验性( _composable ) FSDPv2 的实际配置示例展示基于 DeviceMesh 与 mixed precision 的逐层分片组合:
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed._composable.fsdp import fully_shard, MixedPrecisionPolicy
dp_size = int(os.environ["WORLD_SIZE"])
mesh = init_device_mesh("cuda", (dp_size,), mesh_dim_names=("dp",))
mp_policy = MixedPrecisionPolicy(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
cast_forward_inputs=True,)
Per-layer FSDP wrapping
for layer_id, transformer_block in enumerate(model.layers): fully_shard( transformer_block,
mesh=mesh,
mp_policy=mp_policy,
reshard_after_forward=True, # Same as FSDP1 forward_prefetch
reshard_after_backward=True, # Controls prefetch behavior) fully_shard(model, mesh=mesh, mp_policy=mp_policy) FSDPv2 的 reshard_after_forward 等价于 FSDP1 的 forward_prefetch=False (设置为 True 等价于最小显存模 式), reshard_after_backward 控制梯度是否立即 reduce-scatter。与 FSDP1 相比,FSDPv2 不依赖全局 auto_wrap_policy ,而是让用户显式标注每个模块的分片策略——这种声明式 API 在复杂模型(如 MoE、多模态)中 更可控。 参考文献:Rajbhandari et al., “ZeRO: Memory Optimizations Toward Training Trillion Parameter Models”, SC 2020。Zhao et al., “PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel”, VLDB 2023。
6.4 张量并行 Megatron-LM
张量并行(Tensor Parallelism, TP)是将单个 Transformer 层的权重矩阵按维度切分到多个 GPU,每步前向/反向传播 中插入通信操作。Megatron-LM(Shoeybi et al., 2019)首次系统化地将 TP 应用于 Transformer 结构,现已成为大模型 训练的标配。
6.4.1 列并行线性层
对权重矩阵 A ∈ R ,将 A 按列切分为 [A , A , …, A ],每个子矩阵 A ∈ R k×m 1 2 k×m/tp 位于不同 GPU,其前向计算流程如图 6-4所示。 tp i Column-Parallel Input X (n, k) broadcast / same input GPU0: A1 GPU1: A2 GPU2: A3 (k, m/tp) (k, m/tp) (k, m/tp) Y1 = X.A1 Y2 = X.A2 Y3 = X.A3 (n, m/tp) f function: All-Gather Full output Y (n, m) 图6-4 列并行线性层的前向计算流程 •前向传播:输入 X 在各 GPU 上相同(不切分),各自计算 Y = XA ,输出 Y 保持切分状态直接传入后续行并行层 ( f 函数 forward=Identity,无需 All-Gather);若下游需要完整输出(如接 Softmax),才通过 All-Gather 拼接得到 i i i 完整 Y 。 •反向传播:输出梯度 在各 GPU 上已切分,直接计算权重梯度 = X ⋅ ;输入梯度 = ∑ A 需要 All- ∂L ∂L T ∂L ∂L ∂L T Reduce( f 函数 backward=All-Reduce)。 ∂Yi ∂Ai ∂Yi ∂X i ∂Yi i
6.4.2 行并行线性层
权重 B ∈ R 按行切分为 [B , B , …, B ] : m×n T T T T tp •前向传播:输入 X 也按列切分(由前面列并行层的输出自然产生),计算 Z = X B ,然后 All-Reduce 得到完整 Z 。 i i i •反向传播:Z 的梯度在各 GPU 上相同(因为前向做了 All-Reduce),各 GPU 独立计算 B 和 X 的梯度,X 的梯度通过 All-Gather( g 函数)汇聚。 i i i
6.4.3 通信模式总结
每层 Transformer 共 4 次 All-Reduce:前向的 Attention 输出投影与 MLP down 投影各一次(行并行),反向的 QKV 与 MLP gate/up 输入梯度各一次(列并行)。列并行在反向对输入梯度做 All-Reduce,行并行在前向对输出做 All-Reduce, 二者恰好互补,对应 Megatron 的 f/g 算子,满足自动微分的对偶性质。
6.4.4 MLP 张量并行
Transformer 的 MLP 层天然适配列并行 + 行并行的组合:
# Megatron-LM MLP pseudocode
class ParallelMLP(nn.Module):
def __init__(self, hidden_size, ffn_size, tp_size):
# Column-parallel: W_gate, W_up in R^(hidden, ffn_size/tp)
self.W_gate = ColumnParallelLinear(hidden_size, ffn_size // tp_size)
self.W_up = ColumnParallelLinear(hidden_size, ffn_size // tp_size)
# Row-parallel: W_down in R^(ffn_size/tp, hidden)
self.W_down = RowParallelLinear(ffn_size // tp_size, hidden_size)
def forward(self, x):
# x in R^(b*s, hidden), same on each GPU
gate = gelu(self.W_gate(x)) # output: (b*s, ffn/tp)
up = self.W_up(x) # output: (b*s, ffn/tp)
intermediate = gate * up # element-wise: (b*s, ffn/tp)
output = self.W_down(intermediate) # AllReduce + (b*s, hidden)
return outputTP 的 MLP 将 FFN 矩阵按列切分( W_gate , W_up ),中间激活留在切分的 GPU 上,最后通过行并行 W_down 的 All- Reduce 恢复完整输出。
6.4.5 Attention 张量并行
Multi-Head Attention 天然按头(head)切分:
class ParallelAttention(nn.Module):
def __init__(self, hidden_size, num_heads, tp_size):
self.num_heads_per_gpu = num_heads // tp_size
# QKV projection: column-parallel
self.W_qkv = ColumnParallelLinear(hidden_size, 3 * self.num_heads_per_gpu * head_dim )
# Output projection: row-parallel
self.W_out = RowParallelLinear(
self.num_heads_per_gpu * head_dim, hidden_size)
def forward(self, x):
# x: (b*s, hidden)
qkv = self.W_qkv(x) # (b*s, 3 * heads_per_gpu * head_dim)q, k, v = qkv.chunk(3, dim=-1)
# Each GPU independently computes attention for its heads
attn_out = self_attention(q, k, v) # (b*s, heads_per_gpu * head_dim)
output = self.W_out(attn_out) # AllReduce
return output在 Multi-Head Attention 中每个 head 独立计算,不需要中间通信。仅在输出投影(行并行,All-Reduce)时需要通信; QKV 列并行后输出保持切分,各 GPU 独立计算本地 heads 的 Attention。若启用序列并行(SP),All-Reduce 进一步替 换为 Reduce-Scatter。
6.4.6 TP 度约束分析
张量并行度 tp 的首要约束是模型结构的整除性: num_attention_heads % tp == 0 。对于 GQA/MQA 模型约束更严 ——KV heads 更少,实际约束为 num_kv_heads % tp == 0 。 以 Llama-3-8B(num_q_heads=32, num_kv_heads=8)为例:TP 上限为 8,而非 32;若在 16-GPU 节点上训练,超 过 8 的 TP 度无法整除 KV heads,必须引入 Context Parallelism 或 Pipeline Parallel 而非继续提升 TP。这是从 MHA 迁 移到 GQA 架构时常见的配置陷阱:工程师沿用 num_q_heads % tp == 0 检查,却在 KV projection 切分时报错。 此外,FFN 的中间隐藏维度(intermediate size)也须能被 tp 整除。实践中 TP 通常取 2、4、8。
6.4.7 TP 度选择决策树
选择 TP 度的决策流程如图6-5所示: Start: Model & GPU config num_kv_heads % tp == 0? No Yes Reduce tp or use CP hidden_size < 4096? Yes No tp=1-2: small model model params > 70B? Yes No tp=8: maximize node utiliz seq_len > 8192? ation Yes No tp=4-8 + SP: reduce activat tp=2-4: balance compute/c ion omm 图6-5 TP 度选择决策树 核心原则: •GQA 约束:以 KV heads 数为 TP 上限,而非 Q heads。 •小模型收益递减:隐藏维度 < 4096 时通信占比过高。 •长序列组合:大模型 + 长序列场景优先考虑 TP+SP 组合。 •节点边界:单节点内 TP 不超过 NVSwitch domain size(通常为 8)。 参考文献:Shoeybi et al., “Megatron-LM: Training Multi-Billion Parameter Language Models Using Model Parallelism”, arXiv 2019。Narayanan et al., “Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM”, SC 2021。
6.5 张量并行数学分析
本节从数学上严格推导张量并行的前向/反向传播,并从通信量角度分析 TP 的最优配置。
6.5.1 列并行数学推导
考虑列并行线性层 Y = XA,其中 A ∈ R 被按列切分为 A = [A ∣A ∣ ⋯ ∣A ],A ∈ R 。输入 X ∈ R 在各 k×m 1 2 k×m/tp bs×k GPU 上相同。 tp i 前向传播:各 GPU 计算 Y = XA ∈ R 。完整输出 Y = [Y ∣Y ∣ ⋯ ∣Y ] 在非 Identity 优化时通过 All-Gather 获得, bs×m/tp 1 2 在标准 TP(输出直接进入行并行层)时 f 为 Identity,无需通信。 i i tp 令 f 为 Identity 函数(前向无操作,反向传播做 All-Reduce)。实际 Megatron-LM 实现中的 f 函数: •Forward: All-Gather(或 Identity,如果需要) •Backward: 对于 = ∑ A ,实际实现中通过 AllReduce 在反向传播时完成。 ∂L ∂X ∂L i ∂Yi T i 反向传播:令 ∈ R 为完整输出梯度。 ∂L ∂Y bs×m 权重梯度: ∂L ∂L = XT ⋅ ∈ Rk×m/tp ∂Ai ∂Yi 输入梯度(需要通信): tp ∂L ∂L T =∑ A ∂X ∂Yi i i=1 由于每个 GPU 只能计算自己的局部贡献 ∂L T ∂Yi Ai ,需要通过 All-Reduce 求和获得完整 。 ∂L ∂X
6.5.2 行并行数学推导
B1 行并行线性层 Z = Y B,其中 B ∈ R 按行切分为 B = m×n B2 ⋮ ,B ∈ R i m/tp×n 。 Btp 前向传播:输入 Y 被自然按列切分(来自前面列并行层的输出),各 GPU 计算: Zi = Yi Bi ∈ Rbs×n 完整输出 Z = ∑ Z 通过 All-Reduce 获得。 tp i=1 i 反向传播:完整输出梯度 在各 GPU 上相同(因为前向做了 All-Reduce)。 ∂L ∂Z 权重梯度: ∂L ∂L = YiT ⋅ ∈ Rm/tp×n ∂Bi ∂Z 输入梯度(需要通信): ∂L ∂L T = B ∈ Rbs×m/tp ∂Yi ∂Z i ∂ ∂L 完整 ∂Y1 通过 All-Gather 获得。Megatron-LM 称之为 g 函数。 ∂L = ⋯ ∂Y ∂L ∂Ytp
6.5.3 通信量分析
对一层 Transformer(MLP + Attention),设 batch size b,序列长度 s,隐藏维度 h。每层共 4 次 All-Reduce:前向的 Attention 输出投影(行并行)与 MLP down 投影(行并行)各一次,反向的 QKV 输入梯度与 MLP gate/up 输入梯度 (列并行)各一次。 按环形 AllReduce 复杂度模型,单次 All-Reduce 每 GPU 的环上流量为 2bsh(P − 1)/P ,P ≥ 8 时近似为 2bsh。因此单层 每 GPU 总通信量约 8bsh 字节: (tp) P −1 Tcomm = 4 ⋅ 2bsh ⋅ ≈ 8bsh P 若启用序列并行(SP),All-Reduce 被 All-Gather 与 Reduce-Scatter 替换,通信总量基本不变,收益体现在激活值显存 的降低。
6.5.4 通信计算比建模
- 通信计算比推导 TP 中每 GPU 的计算量(FLOPs)为: Fper GPU = tp ⋅ (24bsh2 + 4bs2 h) (仅MLP+Attention的矩阵乘法) 通信-计算比: 8 ⋅ bsh/B 8 ⋅ C ⋅ tp RTP = =
Fper GPU /C B ⋅ (24h + 4s) 其中 B 为 GPU 间带宽,C 为计算峰值。当 R > 1 时通信是瓶颈。 TP 2) 通信时间细粒度建模 上述 R 将通信视为带宽主导的批量传输。实际 NCCL 通信遵循 α-β 模型: TP Tcomm = α + β ⋅ M 其中 α 为启动延迟(NVSwitch 下约 2-5 μs),β = 1/B 为带宽倒数,M 为单次消息大小。按环形 AllReduce 复杂度模型 直接代入:单次 All-Reduce 每 GPU 环上流量为 2M (P − 1)/P ,延迟为 2(P − 1)(α + βM /P )。 以 Llama-7B(h = 4096,b = 1,s = 2048,bf16)在 H100 NVSwitch 为例:单次 All-Reduce 消息约 2bsh = 33.5 MB, 每层 4 次合计 8bsh = 134 MB,带宽项 134/450 ≈ 0.30 ms。结论:高带宽低延迟的 NVSwitch 下,几十 MB 级别的消息完 全由带宽主导,α 项可忽略。 但当序列较短(s = 512,单次消息约 8 MB)时:带宽项降至约 0.002 ms,α 占比显著上升。这说明小的 TP 通信消息会 使延迟成为瓶颈——这也是 TP 不适合跨节点 IB(α ≈ 10 − 20 μs)的根本原因。 3) 总步时与扩展效率 将计算与通信合并,单步总时间: T1GPU Ttotal = Tcompute + Tcomm = + Tcomm tp 扩展效率(Scaling Efficiency): T1GPU 1 η= = tp ⋅ Ttotal 1 + tp ⋅ Tcomm /T1GPU 当 tp ⋅ T /T > 0.25 时(η < 0.8),TP 的扩展效率已显著衰减。对于 H100 NVSwitch:Llama-7B T ≈ 14 ms/ comm 1GPU 1GPU 层(MLP+Attention),每层通信约 8bsh = 134 MB,带宽项 T ≈ 0.30 ms。tp = 4 时 tp ⋅ T /T = 0.086 → η ≈ comm comm 1GPU 92%,扩展效率优秀;tp = 8 时计算量减半而通信量不变,tp ⋅ T = 0.17 → η ≈ 85%,仍可接受。 /T comm 1GPU TP 前向/反向的完整数据流如图6-6所示。 Backward Forward dL/dZ: (bs, h) Row-Parallel Bwd dL/dY: (bs, m) Column-Parallel Bwd X: (bs, h) Column-Parallel (QKV/W_g Y_i: (bs, m/tp) Row-Parallel (W_out/W_d Z: (bs, h) Replicated dL/dB_i = Y_i^T·dZ All-Gather dL/dY_i Replicated dL/dX = Σ dL/dY_i·A_i^T All-Reduce Σ dL/dX Replicated ate/W_up) Sharded on each GPU own) All-Reduce ΣZ_i Replicated Y_i = X·A_i Z_i = Y_i·B_i 图6-6 TP 前向/反向数据流全景 4) Llama-7B 多 TP 度实验对比 以下为理论建模结果(h = 4096,32 layers,b = 1,s = 2048,bf16,H100 SXM NVSwitch),如表6-5所示。 表6-5 Llama-7B 多 TP 度理论建模 TP度 每GPU参数(M) 每GPU计算(ms/层) 通信量(MB/层) NVLink通信(ms) IB通信(ms) 总时间(ms/层) 扩展效率η 1 6,738 14.0 0 0 0 14.0 1.00 2 3,369 7.0 134 0.30 3.4 7.3 0.96 4 1,685 3.5 134 0.30 3.4 3.8 0.92 8 842 1.75 134 0.30 3.4 2.05 0.85 扩展效率从 TP=2 的 96% 降至 TP=8 的 85%,验证了通信-计算比的恶化趋势。在跨节点 IB(40 GB/s)场景下,TP=2 的 通信时间升至 3.4 ms,总时间 10.4 ms(η = 0.67),此时 TP 已不可行,必须切换为 PP 或 DP 策略。
6.5.5 TP 度最优选择
给定 N 个 GPU 训练一个模型,TP 度选择需考虑: •节点内 vs 节点间:节点内 NVLink(A100:600 GB/s,H100:900 GB/s)适合 TP;节点间 IB/RoCE(约 50 GB/s per link)不适合高频的 TP 通信。 •最优 TP 度:对于隐藏维度 h 的模型和节点内 GPU 数 G : node h tpopt = min Gnode , 8bsB C 实践中,TP=8 是 NVSwitch 节点的典型选择,模型较小时 TP=4 或 2。
6.5.6 激活检查点影响
激活检查点与 TP 结合时,checkpoint 区域内的操作被重新计算,导致额外的通信。以 Transformer layer 为 checkpoint 粒度: •无 checkpoint:每层 8bsh 通信/GPU •有 checkpoint:前向时通信仍为 8bsh/GPU,但反向重计算前向时再次通信,总计 16bsh/GPU 因此,在 TP 场景下 checkpoint 会加倍通信量——这是将 TP+PP 组合时倾向于由 PP 阶段做 checkpoint 的原因之一。
6.5.7 并行策略通信对比
各并行策略的通信特征如表6-6所示。 表6-6 并行策略通信对比 并行策略 通信频率 每次通信量 适合拓扑 TP 每层多次 小(O(bsh)) 节点内(NVLink) PP 每 micro-batch 1-2 次 中(O(bsh)) 节点间 IB 低带宽 DP 每步 1 次 大(O(P ) 参数) 节点间 IB FSDP 每层 1 次 中(O(P ) 参数/GPU) 节点间 IB 这四种并行策略的组合(3D/5D Parallelism)是训练千亿以上参数模型的基础方法论。 参考文献:Narayanan et al., “Efficient Large-Scale Language Model Training on GPU Clusters Using Megatron-LM”, SC 2021。Korthikanti et al., “Reducing Activation Recomputation in Large Transformer Models”, MLSys 2023。
6.6 Megatron 序列并行基础
当序列长度 s 增大时,Transformer 层的激活值内存占用急剧增长。标准 TP 中 LayerNorm/RMSNorm 和 Dropout 层的 输入 X ∈ R 在所有 TP rank 上完全相同,即这些“非参数化”操作在各 GPU 上做了冗余计算和冗余存储。序列并行 b×s×h (Sequence Parallelism, SP)的目标是将这些冗余操作也分布到 TP 的各个 GPU 上。
6.6.1 核心思想与流程
Megatron-LM v3(Korthikanti et al., 2023)提出:在 Transformer 层中,将输入 X 的序列维度(sequence dimension)也按照 TP 度切分,仅在需要完整序列的操作(如 Attention)前后通过通信恢复。标准 TP 与序列并行的对 比如图6-7所示。 SeqParallel StdTP Input X: (b, s, h) LN/Dropout Input X: (b, s, h) LN/Dropout Sharded along s dim to s/t Compute separately All-Gather (s dim) Attention (TP) Reduce-Scatter (s dim) LN/Dropout Each GPU has full s Redundant compute Attention (TP) p 图6-7 序列并行在 Transformer 层中的集成
6.6.2 LayerNorm 与 RMSNorm
归一化层沿隐藏维度规约,与序列维度 s 无关,因此序列切分不改变其数值结果,SP 可以直接嵌入 Transformer 层而不 改变 LayerNorm/RMSNorm 的语义。LayerNorm 计算涉及整个序列的均值和方差: h h xij − μi 1 1 ^ij = , μi = ∑ xij , σi = ∑(xij − μi )2 x σi2 + ϵ h j=1 h j=1 注意 μ 和 σ 是对隐藏维度 h 求的,与序列长度 s 无关,因此 LayerNorm 在 s 维度上是完全独立的。如果序列按 s 维度 切分,各 GPU 可以独立计算 LayerNorm,无需通信。 i i RMSNorm(Llama 系列的标准归一化)计算 x^ = x / ∑ x + ϵ,规约同样沿隐藏维度 h 而非序列维度 s,因此在 s i i 1 h 2 维度上完全独立,同样支持 SP。 h j=1 j
6.6.3 通信模式详解
在 Transformer 层中,SP 引入的通信操作: Forward: Sub-Layer Input (b, s/tp, h) ▶ LayerNorm/RMSNorm (local, no comm)
-> All-Gather on seq dim: (b, s, h) <- Restore full sequence
-> Attention (TP, has AllReduce) <- Standard TP forward
-> Reduce-Scatter on seq dim: (b, s/tp, h) <- Back to sequence parallel▶ Dropout (local, no comm) ▶ residual add (local) Sub-Layer Input (b, s/tp, h) ▶ LayerNorm/RMSNorm (local, no comm) ▶ All-Gather on seq dim: (b, s, h) -> MLP (TP, has AllReduce) <- Standard TP forward ▶ Reduce-Scatter on seq dim: (b, s/tp, h) ▶ Dropout (local, no comm) ▶ residual add (local)
6.6.4 通信量对比
设 TP 度为 p,每层 Transformer 的激活值 X ∈ R 。SP 将 Attention 与 MLP 输出处的一部分 All-Reduce 替换为 All- b×s×h Gather + Reduce-Scatter 的组合,通信量对比如表6-7所示。 表6-7 TP 与 SP 通信量对比 方案 通信原语 每层总通信量 TP only All-Reduce ×4 约 8bsh TP + SP All-Reduce + All-Gather + Reduce-Scatter 约 8bsh Reduce-Scatter 只向每个 rank 写入 1/p 的数据,与新增的 All-Gather 相抵后,SP 的通信总量与 TP 基本不变(约 8bsh );真正的收益是 LayerNorm/Dropout 激活值显存降至 1/p。
6.6.5 激活值内存节省
这是 SP 最直接的优势。每 Transformer 层的激活值(不含 attention 中间结果)约 bsh 个元素(FP16 下 2bsh 字节)。在 TP 度 p 下: •TP only:每 GPU 存储完整激活,2bsh 字节 •TP + SP:每 GPU 存储 2bsh/p 字节 对于 Llama-70B(h = 8192),b = 1, s = 4096,一层激活值 64MB(FP16)。8 路 TP 下 SP 可为一层节省约 56MB。累计 80 层,节省约 4.5GB 激活值内存。这在边界内存情况下至关重要。
6.6.6 FlashAttention 结合
FlashAttention 分块计算 self-attention,天然支持序列维度分片。将 SP 与 FlashAttention 配合: •SP 的 All-Gather 将切分的序列值恢复为完整序列 •FlashAttention 在其分块策略中平滑处理整个序列 •SP 的 Reduce-Scatter 将结果切回分片 三者配合使得长序列训练的内存效率最优。Megatron-LM 最新版本已默认开启 SP 模式。
6.6.7 SP+TP 实现细节
SP 最关键的影响是改变 TP 内部的通信模式。原本的 TP Attention(行并行输出)需要进行 All-Reduce 来聚合各个 TP rank 的 Z 得到完整输出 Z 。启用 SP 后,由于每个 rank 只需要其序列片段,All-Reduce 被替换为一对 All-Gather + Reduce-Scatter: i TP only (Attention output): Z_i (b, s, h) on each rank ▶ AllReduce ▶ Z (b, s, h) replicated TP + SP (Attention output): Z_i (b, s/tp, h) on each rank ▶ AllGather (s dim) ▶ Z_full (b, s, h) ▶ Attention with full s ▶ ReduceScatter (s dim) ▶ Z_i (b, s/tp, h) SP+TP 在 Transformer 层中的完整数据流如图6-8所示。 Rank 0 (s=0..s/tp) Rank 1 (s/tp..2s/tp) Rank 2 (2s/tp..3s/tp) Rank 3 (3s/tp..s) Sub-Layer Input (b, s/tp, h) RMSNorm (local on s/tp) RMSNorm (local on s/tp) All-Gather along seq dim Gather s fragments ▶ full (b, s, h) Attention / MLP with TP TP QKV projection (col-parallel) Local heads attention TP output projection (row-parallel, now uses Reduce-Scatter) Reduce-Scatter along seq dim Scatter back to (b, s/tp, h) Dropout + residual (local) Dropout + residual (local) Rank 0 (s=0..s/tp) Rank 1 (s/tp..2s/tp) Rank 2 (2s/tp..3s/tp) Rank 3 (3s/tp..s) 图6-8 SP+TP 在 Transformer Layer 中的完整数据流
6.6.8 Megatron-LM 配置
环境要求:Python 3.8+,Megatron-LM(对应 PyTorch 2.x)。
# Megatron-LM sequence parallel configuration
from megatron.core.tensor_parallel import model_parallel_cuda_manual_seed
from megatron.core.parallel_state import initialize_model_parallel
# TP=8initialize_model_parallel(
tensor_model_parallel_size=8,
pipeline_model_parallel_size=1,
sequence_parallel=True, # Enable SP)
Transformer Layer auto-inserts All-Gather / Reduce-Scatter pairs
对于长序列(s ≥ 8192)训练,序列并行的内存节省效果尤为显著。 参考文献:Korthikanti et al., “Reducing Activation Recomputation in Large Transformer Models”, MLSys 2023。Jacobs et al., “Sequence Parallelism: Long Sequence Training from System Perspective”, ACL 2023。
6.7 混合并行自动切分算法
随着模型规模和 GPU 集群规模的增长,手工配置 DP + TP + PP 的组合参数(并行度分配、设备网格布局、参数分片策 略)变得越来越困难。自动并行化(Auto-Parallelization)旨在通过编译优化和成本模型,自动搜索最优的并行策略。
6.7.1 自动并行分类法
自动并行化方法可分为三个层次,如图6-9所示: Auto-Parallelization Inter-Operator Parallelism Intra-Operator Parallelism Unified Search Device Placement Pipeline Partitioning Data Parallel Tensor Parallel Dynamic Programming Reinforcement Learning (FlexFlow, GSPMD) (PipeDream, Alpa) (SPMD) (Mesh-TensorFlow) (Alpa, Unity) (REINFORCE-based) 图6-9 自动并行化方法分类
6.7.2 Alpa 两层编译器
Alpa(Zheng et al., OSDI 2022)是第一个将算子间并行和算子内并行统一优化的编译器。其核心贡献是两层优化架构: •算子内并行层(Intra-op pass):以算子为粒度,使用整数线性规划(ILP)确定每个算子的最佳并行策略。定义了一个 设备网格(Device Mesh)概念,∣mesh∣ = tp × dp。 •算子间并行层(Inter-op pass):将计算图划分为多个 Stage(供 PP 使用),使用动态规划最小化端到端延迟。目标函 数: min ∑ max Tstage partition micro-batch stages
6.7.3 设备网格抽象
设备网格是现代自动并行化的核心抽象。一个 2D 网格可以用 (R, C) 表示,总设备数为 R × C :
# PyTorch 2.0+ DeviceMesh API
from torch.distributed.device_mesh import init_device_mesh
# Create 2D mesh: 4-way DP, 2-way TP
mesh = init_device_mesh("cuda", (4, 2), mesh_dim_names=("data_parallel", "tensor_parallel"), # mesh[dp_idx, tp_idx] maps to specific GPUs )
6.7.4 DTensor 与 SPMD
DTensor (Distributed Tensor)是 PyTorch 2.0+ 引入的分布式张量抽象:
from torch.distributed.tensor import DTensor, Shard, Replicate
# Parameter placement on (dp, tp) mesh
spec = (Shard(0), Shard(1)) # shard on both dimensions
dtensor = DTensor.from_local(local_tensor, mesh, spec)
# Common sharding strategy combinations
# (Shard(0), Shard(1)) = fully sharded on both dims
# (Replicate(), Shard(1)) = replicated on dp, sharded on tp
# (Shard(0), Replicate()) = sharded on dp, replicated on tpDTensor 的 redistribute() 方法允许在运行时改变分片布局,自动插入必要的通信操作(All-Gather, Reduce- Scatter 等)。
6.7.5 FlexFlow 搜索算法
FlexFlow(Jia et al., SysML 2019)提出 SOAP(Sample, Operator, Attribute, Parameter)搜索空间:
- Sample:在并行策略空间(SOAP = Sample, Operator, Attribute, Parameter 四个维度)中采样候选配置
- Operate:对采样策略执行模拟执行(simulated execution),计算通信和计算成本
- Average:对多次模拟取平均以降低噪声
- Parallelize:选择成本最小的策略 FlexFlow 的创新在于支持每个算子独立选择并行策略(而非整层统一策略),允许算子间策略不一致时自动插入 tensor redistribution。
6.7.6 手动配置策略
与自动搜索相对,Megatron-LM 采取了一套经过工程验证的手动配置策略。
- Megatron 手工并行配置 Given: N GPUs, model parameters P, memory per GPU M Steps:
- TP = min(8, num_attention_heads) # TP generally does not exceed intra-node GPU count
- Compute model shard params per GPU: P / (TP × PP)
- Determine PP from memory constraint: P / TP / PP × (param + grad + optimizer) ≤ M
- DP = N / (TP × PP) # remaining GPUs for Data Parallel 这套手动策略虽然简单,但在实践中高效——Megatron-LM 论文展示了在 2,240 块 A100 上训练 530B 参数模型(TP=8, PP=35, DP=8)的成功案例。不过它无法处理 MoE 等结构化稀疏模型。
- PyTorch 流水线并行 API PyTorch 2.3(2024年3月)将原 PiPPy 项目正式并入主仓库,以 torch.distributed.pipelining 提供开箱即用的流 水线并行 API。核心抽象是 SplitPoint (标注模型切分位置)和 Schedule (选择 1F1B 或 GPipe 调度):
from torch.distributed.pipelining import pipeline, SplitPoint, Schedule1F1B
# Annotate split points without modifying the model
model_spec = pipeline(model, mb_args=(example_micro_batch,), split_spec={ "transformer.layers.4": SplitPoint.BEGINNING, "transformer.layers.8": SplitPoint.BEGINNING, "transformer.layers.12": SplitPoint.BEGINNING, }, )
# Build per-rank stage from pipeline spec
stage = model_spec.build_stage(stage_idx=rank, device=device)
# 1F1B schedule with n_microbatches
schedule = Schedule1F1B(stage, n_microbatches=8)
schedule.step(micro_batch)与 Megatron 手写方式相比,该 API 通过 torch.fx 自动追踪计算图、内置 1F1B 和 interleaved 1F1B 调度,并与 torch.compile 兼容(可对每 stage 独立编译)。实践限制:含 Python 控制流的动态模型需手动标注切分点;MoE 路 由与流水线气泡的交互需额外处理。
6.7.7 成本模型
自动搜索需要一个快速评估策略好坏的成本模型。通常包含: •计算成本:FLOPs 总数 / 每 GPU 的计算峰值 × 利用率 •通信成本:对每种通信(AllReduce/AllGather/ReduceScatter/AllToAll),使用 α-β 模型估算时间 •内存成本:参数 + 梯度 + 优化器状态 + 激活值 一个典型的复合成本函数: C(strategy) = w1 ⋅ Tcompute + w2 ⋅ Tcomm + w3 ⋅ max(0, Mpeak − Mlimit )2
6.7.8 实际挑战
自动并行化的核心挑战在于搜索空间爆炸:对于有 L 层的模型和 k 种可能的并行配置,搜索空间为 k 。实践中还需要考 L 虑: •算子融合后的新粒度 •TP 通信不能在节点间频繁发生的拓扑约束 •流水线并行的气泡计算 •混合精度下的通信精度转换 PyTorch 的 torch.distributed.tensor.parallel 模块提供了基于 DTensor 的张量并行 API,但自动策略搜索(如 torch.distributed._tensor.auto_parallel )仍在快速发展中。 参考文献:Zheng et al., “Alpa: Automating Inter- and Intra-Operator Parallelism for Distributed Deep Learning”, OSDI 2022。Jia et al., “Beyond Data and Model Parallelism for Deep Neural Networks”, SysML 2019。Unger et al., “Unity: Accelerating DNN Training Through Joint Optimization of Algebraic Transformations and Parallelization”, OSDI 2022。
6.8 FSDP 生产配置与调优
本节以 Llama-7B 模型在 8×A100(80GB)上的 FSDP 训练为场景,提供完整的生产级配置和性能调优指南。
6.8.1 硬件环境与脚本
- 硬件环境 •GPU:8× NVIDIA A100-SXM4-80GB(NVLink 互联,600 GB/s) •网络:Mellanox ConnectX-6 HDR 200Gb/s InfiniBand •PyTorch 2.2+,CUDA 12.1,NCCL 2.19+
- 训练脚本 环境要求:Python 3.10+,PyTorch 2.2+,CUDA 12.1。
import os
import contextlib
import functools
import torch
import torch.distributed as dist
from torch.distributed.fsdp import (FullyShardedDataParallel as FSDP, MixedPrecision, CPUOffload, ShardingStrategy, BackwardPrefetch, StateDictType ) from torch.distributed.fsdp.wrap import ( transformer_auto_wrap_policy, size_based_auto_wrap_policy, ) from torch.distributed.fsdp.fully_sharded_data_parallel import ( FullOptimStateDictConfig, FullStateDictConfig ) from torch.distributed.algorithms._checkpoint.checkpoint_wrapper import ( checkpoint_wrapper, CheckpointImpl, apply_activation_checkpointing ) import torch.distributed._shard.checkpoint as dist_cp from transformers import ( AutoModelForCausalLM, AutoTokenizer, AutoConfig, get_cosine_schedule_with_warmup )
from transformers.models.llama.modeling_llama import LlamaDecoderLayer
def init_distributed():
dist.init_process_group(backend='nccl')
local_rank = int(os.environ['LOCAL_RANK'])
torch.cuda.set_device(local_rank)
return local_rank, dist.get_rank(), dist.get_world_size()
# ==========
def get_fsdp_config():
return {'sharding_strategy': ShardingStrategy.FULL_SHARD, 'cpu_offload': CPUOffload(offload_params=False), 'mixed_precision': MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.float32,
buffer_dtype=torch.bfloat16,
keep_low_precision_grads=False,), 'backward_prefetch': BackwardPrefetch.BACKWARD_PRE, 'forward_prefetch': True, 'limit_all_gathers': True, 'use_orig_params': True, 'sync_module_states': True, 'device_id': torch.cuda.current_device(),
}
# ==========
def get_auto_wrap_policy():
transformer_auto_wrapper = functools.partial(transformer_auto_wrap_policy, transformer_layer_cls={LlamaDecoderLayer}, )
return transformer_auto_wrapper
# ==========
def apply_fsdp_checkpointing(model):
non_reentrant_wrapper = functools.partial(checkpoint_wrapper, checkpoint_impl=CheckpointImpl.NO_REENTRANT, ) apply_activation_checkpointing( model, checkpoint_wrapper_fn=non_reentrant_wrapper, check_fn=lambda submodule: isinstance(submodule, LlamaDecoderLayer) ) def main(): local_rank, rank, world_size = init_distributed()
# Load model and tokenizer
model_name = "meta-llama/Llama-2-7b-hf"
config = AutoConfig.from_pretrained(model_name)
model = AutoModelForCausalLM.from_pretrained(model_name,
config=config,
torch_dtype=torch.bfloat16,
device_map=None, # FSDP manages device placement)
# Apply activation checkpointing
apply_fsdp_checkpointing(model)
# Apply FSDP wrapping
fsdp_config = get_fsdp_config()
model = FSDP(model, auto_wrap_policy=get_auto_wrap_policy(), **fsdp_config )
# Optimizer
optimizer = torch.optim.AdamW(
model.parameters(),
lr=3e-4,
betas=(0.9, 0.95),
weight_decay=0.1,
fused=True, # CUDA fused AdamW) # Learning rate scheduler scheduler = get_cosine_schedule_with_warmup( optimizer, num_warmup_steps=2000, num_training_steps=50000 )
# ========== Training Loop ==========
global_batch_size = 128 # equivalent large batch
micro_batch_size = 2 # per GPU per step
grad_accum_steps = global_batch_size // (micro_batch_size * world_size)
model.train()
for step, batch in enumerate(train_dataloader):
# Gradient accumulation
for i, micro_batch in enumerate(batch_iter(batch, micro_batch_size)):
is_last_micro_step = (i == grad_accum_steps - 1)
# Only trigger AllReduce on last stepwith (model.no_sync() if not is_last_micro_step else contextlib.nullcontext()):
loss = model(**micro_batch).loss / grad_accum_steps
loss.backward()
# Gradient clipping (FSDP auto-handles sharded gradients)
grad_norm = model.clip_grad_norm_(max_norm=1.0)
optimizer.step()
scheduler.step()
optimizer.zero_grad(set_to_none=True)
# Periodically save checkpoints
if step % 1000 == 0 and step > 0:with FSDP.state_dict_type( model, StateDictType.FULL_STATE_DICT, FullStateDictConfig(offload_to_cpu=True, rank0_only=True), FullOptimStateDictConfig(offload_to_cpu=True, rank0_only=True), ):
state_dict = model.state_dict()
opt_state = FSDP.optim_state_dict(model, optimizer)
if rank == 0:
torch.save({'model': state_dict, 'optimizer': opt_state, 'step': step, }, f'checkpoint-step-{step}.pt')
6.8.2 关键参数解析
if rank == 0 and step % 10 == 0:
print(f"Step {step}: loss={loss.item():.4f}, grad_norm={grad_norm:.2f}")
dist.destroy_process_group()- backward_prefetch • BACKWARD_PRE (推荐):在当前层 backward 期间异步预取下一层的参数。 • BACKWARD_POST :在当前层 backward 完成后才预取下一层(通信重叠少)。 •实测:BACKWARD_PRE 可使训练吞吐提升 10-15%。
- limit_all_gathers • True :限制同时排队的 All-Gather 操作数。 •避免一次性发起大量通信挤爆 NCCL 内部缓冲区。 •在 8 GPU 规模下通常不会成为瓶颈。
- forward_prefetch • True :Forward 阶段预取下一层的 All-Gathered 参数。 •对于计算量较小的模型层(如小 hidden_size),Forward Prefetch 效果不明显,甚至可能因通信竞争而退化。
6.8.3 性能测试结果
在 8×A100(80GB)上训练 Llama-7B,不同配置的吞吐对比如表6-8所示。 表6-8 不同配置的训练吞吐对比 配置 每GPU micro_batch 吞吐 (tokens/s) 显存使用 (GB) MFU DDP (无FSDP) 2 28,000 72.5 52% FSDP Full Shard 2 26,500 38.2 49%
FSDP + Activation Ckpt 2 24,200 28.7 45%
FSDP + Ckpt + bf16 4 30,100 42.1 55%
FSDP + Ckpt + bf16 + grad_accum=4 4 31,400 42.8 58%MFU 提升策略: •bfloat16 比 FP16 有更快的数值转换(无需 loss scaling) • fused=True 的 AdamW 减少 kernel launch 开销 •梯度累积增加有效 batch size,提升 GPU 计算占空比
6.8.4 Profiling 方法
# Use PyTorch Profiler to analyze
python -m torch.distributed.run --nproc_per_node=8 train.py \
--profile-with-tensorboard
# NCCL env vars for detailed communication loggingexport NCCL_DEBUG=INFO export NCCL_DEBUG_SUBSYS=ALL Profiler 中重点关注的 FSDP 事件: • all_gather 耗时:应该与 backward 有良好重叠 • reduce_scatter 耗时:应与下一层 forward 重叠 • wait_for_gpu 事件:表示通信未完成导致计算停滞 •通信-计算重叠率 = 1 - (wait时间 / 总通信时间)
6.8.5 扩展指南
不同 GPU 规模下的扩展配置如表6-9所示。 表6-9 FSDP 扩展配置指南 GPU 数 每GPU bs 梯度累积步数 等效 batch 最优FSDP配置 1 8 16 128 FSDP无效 2 4 16 128 FULL_SHARD, backward_prefetch=true 4 2 16 128 FULL_SHARD, limit_all_gathers=true 8 2 8 128 FULL_SHARD, forward_prefetch=true
6.8.6 常见问题排查
•OOM 显存溢出:检查 limit_all_gathers 是否开启、激活检查点是否正确应用、 param_dtype 是否是 bfloat16 。 •吞吐不线性扩展:检查通信-计算重叠率( nccl:all_gather 时长应小于层计算时长的 80%)。 •FSDP checkpoint 加载失败:确保保存时使用 FULL_STATE_DICT ,加载时也将模型先包装再加载。 •NCCL timeout:通过 dist.init_process_group(timeout=timedelta(hours=2)) 增大 PyTorch 集合通信超时 (默认1800秒/30分钟,大模型加载阶段可能耗尽);PyTorch 2.1+ 还需设置 TORCH_NCCL_HEARTBEAT_TIMEOUT_SEC=7200 控制 Watchdog 超时( NCCL_TIMEOUT 环境变量已废弃,改由 PyTorch 层统一管理)。 参考文献:PyTorch FSDP 官方文档。Zhao et al., “PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel”, VLDB 2023。