训练框架与运行时

训练框架与运行时

框架的选择不只是技术问题,更是组织能力与生态系统的选择。

章节导入:框架为什么重要

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 通信器
Tip调试 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 output

DeepSpeed: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")
Tip为什么 JAX 值得关注

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 output
Warning编译器的代价

torch.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()
Warning梯度累积的陷阱

使用梯度累积时,如果模型中有 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 三阶段

%%{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

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()
TipFSDP 调优清单
  1. fsdp_group:匹配节点内 NVLink 拓扑(如 8 卡一组)
  2. forward_prefetch=True:在当前层反向传播时预取下一层的参数
  3. limit_all_gathers=True:避免多个 AllGather 同时发起导致 OOM
  4. activation_checkpointing:配合 FSDP 使用,进一步减少激活值内存
  5. 监控 nccl_communicator 队列:如果 AllGather 等待时间长,说明通信是瓶颈
WarningFSDP vs DeepSpeed ZeRO-3

两者原理相似,但实现不同。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)