%%{init: {'theme': 'base'}}%%
graph TB
subgraph "DDP(标准数据并行)"
DDP1["GPU 0: 完整模型 + 完整优化器"]
DDP2["GPU 1: 完整模型 + 完整优化器"]
DDP3["GPU N: 完整模型 + 完整优化器"]
end
subgraph "ZeRO-1 / FSDP-like:优化器分片"
Z1_1["GPU 0: 完整模型 + 优化器 1/N"]
Z1_2["GPU 1: 完整模型 + 优化器 1/N"]
Z1_3["GPU N: 完整模型 + 优化器 1/N"]
end
subgraph "ZeRO-2:优化器 + 梯度分片"
Z2_1["GPU 0: 完整模型 + 梯度 1/N + 优化器 1/N"]
Z2_2["GPU 1: 完整模型 + 梯度 1/N + 优化器 1/N"]
end
subgraph "ZeRO-3 / FSDP:全部分片"
Z3_1["GPU 0: 模型 1/N + 梯度 1/N + 优化器 1/N"]
Z3_2["GPU 1: 模型 1/N + 梯度 1/N + 优化器 1/N"]
Z3_3["GPU N: 模型 1/N + 梯度 1/N + 优化器 1/N"]
end
训练框架与运行时
训练框架与运行时
框架的选择不只是技术问题,更是组织能力与生态系统的选择。
章节导入:框架为什么重要
2017 年,如果你要做分布式训练,选择基本上是 TensorFlow 或 PyTorch。到了 2025 年,选择变成了 PyTorch + DeepSpeed、PyTorch + Megatron-LM、JAX + T5X,或者某个云计算厂商的托管服务。
框架解决的核心问题是:把数学公式变成可以在数千张 GPU 上高效运行的代码。这中间涉及编译优化、内存管理、通信调度、容错恢复等一系列工程挑战。
一个好的训练框架应该让研究人员思考模型架构,而不是思考 AllReduce 的实现细节。
5.1 PyTorch 分布式训练核心机制
PyTorch 已经成为大模型训练的事实标准。理解它的分布式原语是理解所有上层框架的基础。
进程组与后端
import torch.distributed as dist
# 初始化进程组
dist.init_process_group(
backend="nccl", # GPU 通信后端
init_method="env://", # 发现方式
world_size=8, # 总进程数
rank=0 # 当前进程编号
)
# 创建子进程组(用于 3D 并行)
tp_group = dist.new_group(ranks=[0,1,2,3], backend="nccl") # 张量并行组
pp_group = dist.new_group(ranks=[0,4], backend="gloo") # 流水线并行组
dp_group = dist.new_group(ranks=[0,2,4,6], backend="nccl") # 数据并行组核心通信原语
| 原语 | 操作 | 典型用途 |
|---|---|---|
broadcast |
一发多收 | 参数初始化同步 |
all_reduce |
多发多收(聚合) | 梯度同步 |
all_gather |
收集所有分片 | ZeRO-2 梯度收集 |
reduce_scatter |
聚合后分片 | ZeRO-3 梯度分片 |
all_to_all |
全互联 | MoE Token 路由 |
send/recv |
点对点 | 流水线并行 |
# AllReduce 示例:梯度同步的核心
gradient = torch.randn(1024, device="cuda")
dist.all_reduce(gradient, op=dist.ReduceOp.SUM)
gradient /= world_size # 平均梯度
# AllGather 示例:收集各卡的参数分片
param_shard = torch.randn(256, device="cuda") # 本地分片
full_param = [torch.empty(256, device="cuda") for _ in range(world_size)]
dist.all_gather(full_param, param_shard)
# ReduceScatter 示例:ZeRO-3 的核心操作
# 聚合梯度然后分片到各卡
grad_shard = torch.empty(256, device="cuda")
input_tensor = grad_shard.clone()
dist.reduce_scatter(grad_shard, [input_tensor], op=dist.ReduceOp.SUM)ProcessGroupNCCL:底层实现
NCCL(NVIDIA Collective Communications Library)是 GPU 集合通信的底层库。理解它的拓扑感知对性能调优至关重要:
# 控制 NCCL 通信通道
os.environ["NCCL_NET"] = "IB" # 使用 InfiniBand
os.environ["NCCL_IB_HCA"] = "mlx5_0" # 指定网卡
os.environ["NCCL_SOCKET_IFNAME"] = "eth0"
os.environ["NCCL_DEBUG"] = "WARN" # 调试日志级别
os.environ["NCCL_IB_DISABLE"] = "0"
os.environ["NCCL_IB_GID_INDEX"] = "3" # RoCE 配置
# PyTorch 2.x 的延迟初始化
# 只在实际使用时才创建 NCCL 通信器当遇到 NCCL hang 或超时时,第一步是设置 NCCL_DEBUG=INFO。日志会显示每个 rank 的通信通道、网卡选择和拓扑信息。90% 的分布式训练启动问题都是网络配置不正确导致的。
5.2 TensorFlow、JAX、Megatron-LM、DeepSpeed 对比
框架对比一览
| 特性 | PyTorch + DDP/FSDP | PyTorch + Megatron-LM | PyTorch + DeepSpeed | JAX + Flax | TensorFlow |
|---|---|---|---|---|---|
| 心智模型 | 动态图、Eager | 动态图 + 自定义并行 | 动态图 + ZeRO | 函数式、JIT | 静态/动态图 |
| 并行策略 | DDP, FSDP | 3D 并行(TP+PP+DP) | ZeRO 1/2/3 + PP | GSPMD(自动并行) | MirroredStrategy, TPUStrategy |
| 目标用户 | 通用 | 大模型预训练 | 大模型训练 | 研究 + TPU 生产 | 生产部署 |
| TPU 支持 | 部分 | 无 | 部分 | ✅ 原生 | ✅ 原生 |
| 编译优化 | TorchDynamo/Inductor | 部分 | 部分 | XLA | XLA |
| 学习曲线 | 低 → 高(分布式) | 高 | 中 | 中高 | 中 |
Megatron-LM:为 GPT 训练而生
Megatron-LM 是 NVIDIA 开发的大模型训练框架,核心价值是提供了经过深度优化的 3D 并行实现。
# Megatron-LM 的核心组件
from megatron.training import pretrain
from megatron.core import tensor_parallel as tp
from megatron.core.pipeline_parallel import schedule_forward_backward
# 自定义模型需要继承 MegatronModule
class MyGPTModel(MegatronModule):
def __init__(self, config):
super().__init__(config)
# embedding 层(嵌入并行)
self.embedding = VocabParallelEmbedding(
vocab_size, hidden_size
)
# Transformer 层(张量并行 + 流水线并行)
self.decoder = TransformerBlock(
num_layers=num_layers_per_pp_rank,
tensor_model_parallel_size=tp_size,
pipeline_model_parallel_size=pp_size,
)
def forward(self, input_ids):
hidden = self.embedding(input_ids)
output = self.decoder(hidden)
return outputDeepSpeed:ZeRO 的力量
DeepSpeed 的核心创新是 ZeRO(Zero Redundancy Optimizer),它通过将优化器状态、梯度和参数分片到不同 GPU 上来消除数据并行中的内存冗余。
# DeepSpeed 最小示例
import deepspeed
model, optimizer, _, _ = deepspeed.initialize(
model=model,
optimizer=optimizer,
config_params={
"train_micro_batch_size_per_gpu": 4,
"zero_optimization": {
"stage": 2,
"allgather_partitions": True,
"reduce_scatter": True,
"overlap_comm": True,
"contiguous_gradients": True,
"reduce_bucket_size": 5e8,
},
"bf16": {"enabled": True},
"gradient_accumulation_steps": 8,
}
)
# 训练循环和普通 PyTorch 一样
for batch in dataloader:
loss = model(batch)
model.backward(loss)
model.step()JAX:函数式的优雅
import jax
import jax.numpy as jnp
from jax import pmap
# JAX 的并行是函数式的——pmap 自动将函数并行化
def train_step(state, batch):
def loss_fn(params):
logits = model.apply(params, batch)
return cross_entropy(logits, batch.labels)
grads = jax.grad(loss_fn)(state.params)
state = state.apply_gradients(grads=grads)
return state
# pmap:在多设备上并行执行,自动处理梯度同步
parallel_train_step = pmap(train_step, axis_name="batch")
# 所有设备同步梯度(等价于 AllReduce)
grads = jax.lax.pmean(grads, axis_name="batch")JAX 的 GSPMD(Global Single Program, Multiple Device)系统能够自动将单设备程序转换为任意并行策略。你只需要用注释标记张量的分片方式,编译器自动生成通信代码。这是分布式训练的”终极形态”——但它需要你先掌握函数式编程的心智模型。
5.3 计算图优化与编译器
Eager 模式(PyTorch 默认)灵活但慢。编译器可以把动态执行转化为优化的静态图。
TorchDynamo + Inductor(PyTorch 2.x)
import torch
import torch._dynamo
# 一行代码启用编译
model = torch.compile(model, mode="max-autotune")
# PyTorch 2.x 编译管线:
# Python bytecode → TorchDynamo(Guard 机制)
# → FX Graph → AOTAutograd(融合反向传播)
# → Inductor(Triton Kernel 生成)
# → 优化的 GPU Kernel
# 查看编译详情
torch._dynamo.config.verbose = True
torch._dynamo.config.suppress_errors = False
# 控制编译模式
model = torch.compile(model, mode="default") # 快速编译
model = torch.compile(model, mode="reduce-overhead") # 减少 CPU 开销
model = torch.compile(model, mode="max-autotune") # 最大优化XLA:Google 的编译器
import torch_xla
import torch_xla.core.xla_model as xm
# PyTorch/XLA:在 TPU 上运行 PyTorch
device = xm.xla_device()
model = model.to(device)
# 显式编译和执行图
# XLA 通过 Lazy Tensor 机制收集操作,然后一次性编译
for step, batch in enumerate(loader):
loss = model(batch)
loss.backward()
xm.optimizer_step(optimizer) # 触发 XLA 编译 + 执行Triton:自定义 Kernel 的新标准
import triton
import triton.language as tl
@triton.jit
def softmax_kernel(
input_ptr, output_ptr,
input_row_stride, output_row_stride,
n_cols, BLOCK_SIZE: tl.constexpr,
):
row_idx = tl.program_id(0)
row_start = row_idx * input_row_stride
cols = tl.arange(0, BLOCK_SIZE)
mask = cols < n_cols
input = tl.load(input_ptr + row_start + cols, mask=mask, other=-float('inf'))
row_max = tl.max(input, axis=0)
input -= row_max
exp = tl.exp(input)
row_sum = tl.sum(exp, axis=0)
output = exp / row_sum
tl.store(output_ptr + row_idx * output_row_stride + cols, output, mask=mask)
# 使用自定义 Kernel
def softmax(x):
n_rows, n_cols = x.shape
BLOCK_SIZE = triton.next_power_of_2(n_cols)
output = torch.empty_like(x)
softmax_kernel[(n_rows,)](
x, output,
x.stride(0), output.stride(0),
n_cols, BLOCK_SIZE=BLOCK_SIZE
)
return outputtorch.compile 在首次迭代时需要大量的 tracing 和编译时间(可能几分钟到几十分钟)。在调试阶段,建议关闭编译;只在确定模型逻辑正确后开启。另外,某些动态控制流(如 if loss < threshold)可能导致频繁重编译。
5.4 自动微分与梯度累积
Autograd 的工作原理
PyTorch 的 Autograd 在前向传播时自动构建计算图(动态图),在反向传播时沿图的反向计算梯度。
# 理解 Autograd 图
x = torch.randn(3, requires_grad=True)
y = x * 2
z = y.sum()
print(z.grad_fn) # SumBackward0
print(z.grad_fn.next_functions) # ((MulBackward0(...)),)
z.backward()
print(x.grad) # tensor([2., 2., 2.])梯度累积:用小显存模拟大 batch
当 batch 太大放不进显存时,可以将一个大 batch 拆成多个微批次,累积梯度后再更新:
# 梯度累积的标准实现
accumulation_steps = 8
micro_batch_size = 4
# 有效 batch size = 8 * 4 * world_size = 256 (8卡)
optimizer.zero_grad()
for step, batch in enumerate(dataloader):
# 注意:loss 需要除以累积步数
loss = model(batch)
loss = loss / accumulation_steps
# backward 累积梯度
loss.backward()
# 每 accumulation_steps 步执行一次优化器更新
if (step + 1) % accumulation_steps == 0:
# 可选:梯度裁剪
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
optimizer.zero_grad()使用梯度累积时,如果模型中有 BatchNorm 层,统计量(running mean/var)会在每个微批次上更新,而不是在整个 batch 上更新,导致统计偏差。LayerNorm 不受此影响,这也是为什么大模型普遍使用 LayerNorm。
另外,loss / accumulation_steps 很容易被忘记。如果忘记归一化,等效学习率会变成设定值的 \(N\) 倍。
自动微分的进阶技巧
# 梯度检查点(Gradient Checkpointing):用计算换内存
from torch.utils.checkpoint import checkpoint
class CheckpointedTransformer(nn.Module):
def __init__(self, layers):
super().__init__()
self.layers = nn.ModuleList(layers)
def forward(self, x):
for layer in self.layers:
# checkpoint 不保存中间激活值
# 反向传播时重新计算
x = checkpoint(layer, x, use_reentrant=False)
return x
# 混合精度训练下的梯度缩放
scaler = torch.cuda.amp.GradScaler()
with torch.autocast(device_type="cuda", dtype=torch.bfloat16):
loss = model(inputs)
scaler.scale(loss).backward()
scaler.unscale_(optimizer)
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
scaler.step(optimizer)
scaler.update()5.5 FSDP、DDP、ZeRO 内存优化策略
内存消耗的真实数字
以 7B 参数模型为例(AdamW 优化器,BF16 混合精度):
| 组件 | 每参数字节 | 总内存 |
|---|---|---|
| 模型权重(BF16) | 2 | 14 GB |
| 梯度(BF16) | 2 | 14 GB |
| 优化器状态(FP32 momentum + variance) | 8 | 56 GB |
| 模型权重主副本(FP32) | 4 | 28 GB |
| 激活值(seq=2048, batch=8) | ~ | ~20 GB |
| 合计 | ~132 GB |
一张 80GB H100 放不下一个 7B 模型的训练状态!
DDP vs FSDP vs ZeRO 三阶段
ZeRO 各阶段的内存节省
| 阶段 | 分片内容 | 7B 模型每卡内存 (8卡) | 通信开销 |
|---|---|---|---|
| DDP | 无 | ~132 GB(放不下) | 1× AllReduce |
| ZeRO-1 | 优化器状态 | ~82 GB | 1× AllReduce |
| ZeRO-2 | 优化器 + 梯度 | ~68 GB | 1× AllReduce + ReduceScatter |
| ZeRO-3 | 全部 | ~25 GB | AllGather(前向+反向) |
PyTorch FSDP 实战
from torch.distributed.fsdp import (
FullyShardedDataParallel as FSDP,
MixedPrecision,
ShardingStrategy,
CPUOffload,
)
from torch.distributed.fsdp.wrap import (
size_based_auto_wrap,
transformer_auto_wrap_policy,
)
# FSDP 配置
fsdp_config = {
"sharding_strategy": ShardingStrategy.FULL_SHARD, # ZeRO-3
"mixed_precision": MixedPrecision(
param_dtype=torch.bfloat16,
reduce_dtype=torch.bfloat16,
buffer_dtype=torch.bfloat16,
),
"auto_wrap_policy": transformer_auto_wrap_policy, # 按 Transformer 层自动切分
"cpu_offload": CPUOffload(offload_params=False), # 关闭 CPU offload 以保性能
"limit_all_gathers": True, # 限制并发 AllGather 防止 OOM
"forward_prefetch": True, # 预取下一层的参数
"use_orig_params": True, # 保留原始参数名(方便 checkpoint)
}
# 包装模型
model = FSDP(model, **fsdp_config)
# FSDP 训练循环与普通 PyTorch 完全一样
for batch in dataloader:
loss = model(batch)
loss.backward()
optimizer.step()
optimizer.zero_grad()fsdp_group:匹配节点内 NVLink 拓扑(如 8 卡一组)forward_prefetch=True:在当前层反向传播时预取下一层的参数limit_all_gathers=True:避免多个 AllGather 同时发起导致 OOMactivation_checkpointing:配合 FSDP 使用,进一步减少激活值内存- 监控
nccl_communicator队列:如果 AllGather 等待时间长,说明通信是瓶颈
两者原理相似,但实现不同。FSDP 是 PyTorch 原生的,与 PyTorch 生态深度集成;DeepSpeed 提供了更多功能(如 CPU offload、NVMe offload)。如果你追求原生体验和长期维护,选 FSDP;如果你需要极致的内存优化或快速实验,选 DeepSpeed。
FSDP 的 Checkpoint 保存
from torch.distributed.fsdp import FullyShardedDataParallel as FSDP
import torch.distributed.checkpoint as dcp
# 分片保存(每个 rank 只保存自己的分片)
state_dict = {
"model": FSDP.optim_state_dict(model, optimizer),
"step": step,
}
# 使用分布式 checkpoint API(支持异构拓扑恢复)
dcp.save(
state_dict,
checkpoint_id=f"checkpoint_step_{step}",
storage_writer=dcp.FileSystemWriter(f"/checkpoints/step_{step}"),
)
# 恢复
dcp.load(
state_dict,
checkpoint_id=f"checkpoint_step_{step}",
storage_reader=dcp.FileSystemReader(f"/checkpoints/step_{step}"),
)小结
训练框架的选择取决于团队规模、模型类型和基础设施:
- 快速实验 + 中小模型 → PyTorch DDP +
torch.compile - 大模型预训练(GPT/Llama 类) → Megatron-LM 或 DeepSpeed
- 追求极致内存优化 → DeepSpeed ZeRO-3 + CPU/NVMe offload
- TPU 训练或自动并行 → JAX + GSPMD
- 长期维护的开源项目 → PyTorch FSDP(原生支持、社区最大)
最终,最好的框架是你团队能维护和调试的那个。
延伸阅读
- PyTorch FSDP 论文:Zhao et al., “PyTorch FSDP: Experiences on Scaling Fully Sharded Data Parallel” (2023)
- DeepSpeed ZeRO 论文:Rajbhandari et al., “ZeRO: Memory Optimizations Toward Training Trillion Parameter Models” (2020)
- Megatron-LM 系列:Shoeybi et al. (2019); Narayanan et al. (2021)
- GSPMD:Xu et al., “GSPMD: General and Scalable Parallelization for ML Computation Graphs” (2021)
- TorchDynamo:Ansel et al., “TorchDynamo: A Dynamic Compiler for PyTorch” (2023)
- Triton:Tillet et al., “Triton: An Intermediate Language and Compiler for Tiled Neural Network Computations” (2019)