Skip to content

分布式训练 — 从单卡到千卡集群 ​

#分布式训练 · #DeepSpeed · #ZeRO · #FSDP · #数据并行 · #模型并行 · #流水线并行

当单卡放不下模型或训练太慢时,分布式训练是必经之路。本文从数据并行出发,逐步深入到 ZeRO、3D 并行和 FSDP 实战。


为什么需要分布式训练? ​

mermaid
graph TD
    subgraph "单卡训练困境"
        SINGLE["单张 GPU<br/>80GB A100"] --> Q1{"模型大小?"}
        Q1 -->|"< 80GB"| OK["可以训练<br/>但慢"]
        Q1 -->|"> 80GB"| FAIL["❌ OOM<br/>无法训练"]
    end

    subgraph "分布式训练解决"
        DIST["多张 GPU 协作"] --> DP["数据并行<br/>每卡一份模型"]
        DIST --> MP["模型并行<br/>每卡一部分层"]
        DIST --> ZERO["ZeRO 优化<br/>切分优化器状态"]
    end

    style FAIL fill:#e74c3c,color:#fff
    style DIST fill:#2ecc71,color:#fff
模型参数量FP16 权重梯度优化器状态(Adam)总需求能单卡训练?
Qwen4-7B7B14 GB14 GB28 GB56 GB✅ A100 80GB
Qwen4-72B72B144 GB144 GB288 GB576 GB❌ 需要 8×A100
LLaMA 4 Scout109B×MoE218 GB——>1TB❌ 需要多机

注意:上面是"最小需求"。实际训练还需要激活值存储(Activation Memory),通常是权重的 2-5 倍。所以 7B 模型用 LoRA 才能在 24GB 消费级显卡上跑。


数据并行 (Data Parallel, DP) ​

原理 ​

每张 GPU 持有完整模型副本,但处理不同 batch 的数据。梯度在所有 GPU 间平均后统一更新。

mermaid
graph TD
    subgraph "数据并行流程"
        DATA["训练数据<br/>Batch = 256"] --> SPLIT["切分为 4 份"]
        SPLIT --> GPU0["GPU 0<br/>batch 64"]
        SPLIT --> GPU1["GPU 1<br/>batch 64"]
        SPLIT --> GPU2["GPU 2<br/>batch 64"]
        SPLIT --> GPU3["GPU 3<br/>batch 64"]

        GPU0 --> G0["梯度₀"]
        GPU1 --> G1["梯度₁"]
        GPU2 --> G2["梯度₂"]
        GPU3 --> G3["梯度₃"]

        G0 --> ALLREDUCE["AllReduce<br/>平均所有梯度"]
        G1 --> ALLREDUCE
        G2 --> ALLREDUCE
        G3 --> ALLREDUCE

        ALLREDUCE --> UPDATE["每张卡用平均梯度更新参数<br/>保证模型副本一致"]
    end

    style ALLREDUCE fill:#e74c3c,color:#fff

PyTorch DDP 实现 ​

python
"""
PyTorch DistributedDataParallel (DDP) 完整示例
"""
import torch
import torch.nn as nn
import torch.distributed as dist
from torch.nn.parallel import DistributedDataParallel as DDP
from torch.utils.data import DataLoader, DistributedSampler

def setup_distributed(rank: int, world_size: int):
    """初始化分布式环境"""
    dist.init_process_group(
        backend="nccl",         # NVIDIA NCCL 通信库
        init_method="tcp://localhost:12355",
        rank=rank,
        world_size=world_size,
    )
    torch.cuda.set_device(rank)

def cleanup():
    dist.destroy_process_group()

def train_ddp(rank: int, world_size: int):
    """单进程的训练入口"""
    setup_distributed(rank, world_size)

    # 模型 + DDP 包装
    model = nn.Linear(1024, 10).cuda(rank)
    model = DDP(model, device_ids=[rank])

    # 数据集 + 分布式采样器
    dataset = torch.randn(10000, 1024)  # 示例数据
    sampler = DistributedSampler(
        dataset,
        num_replicas=world_size,
        rank=rank,
        shuffle=True,
    )
    dataloader = DataLoader(dataset, batch_size=64, sampler=sampler)

    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)
    criterion = nn.MSELoss()

    for epoch in range(10):
        sampler.set_epoch(epoch)  # 每 epoch 重新 shuffle
        for batch in dataloader:
            data = batch.cuda(rank)
            target = torch.randn(64, 10).cuda(rank)

            output = model(data)
            loss = criterion(output, target)

            optimizer.zero_grad()
            loss.backward()       # DDP 自动 AllReduce 梯度
            optimizer.step()

        if rank == 0:
            print(f"Epoch {epoch}, Loss: {loss.item():.4f}")

    cleanup()

# 启动方式 (命令行):
# torchrun --nproc_per_node=4 train.py

数据并行的瓶颈 ​

问题原因影响
显存冗余每张卡存完整优化器状态训练大模型受单卡显存限制
通信开销AllReduce 传输整个梯度卡数越多通信占比越大
负载不均最后一张卡可能数据少计算完成时间不同步

通信量量化 ​

数据并行的通信瓶颈可以用公式精确定量。对于参数量为 Φ 的模型:

通信原语单步通信量 (每卡)说明
DDP AllReduce2⋅Φ⋅N−1N bytes梯度总大小 × 2(Ring发送+接收),N 为卡数
ZeRO-3 AllGatherΦ⋅N−1N bytes前向时收集他卡参数
ZeRO-3 ReduceScatterΦ⋅N−1N bytes反向时归约梯度

实例:7B 模型 FP16 训练,Φ = 14 GB,64 张卡下:

  • DDP 每步通信量 ≈ 2×14×63/64 ≈ 27.6 GB/卡——这就是为什么纯数据并行扩展性差
  • ZeRO-3 每步(前向+反向)≈ 2×14×63/64 ≈ 27.6 GB/卡——与 DDP 总量相同,但显存优势巨大
python
"""
通信量计算器
"""
def communication_cost(
    params_gb: float,      # 参数量 (GB, FP16)
    num_gpus: int,         # GPU 数量
    strategy: str = "ddp", # ddp | zero3
) -> dict:
    """计算单步通信量"""
    phi = params_gb * 1e9  # 转为 bytes
    factor = (num_gpus - 1) / num_gpus

    if strategy == "ddp":
        # Ring AllReduce: 2× 数据量
        per_gpu = 2 * phi * factor
    elif strategy == "zero3":
        # AllGather(fwd) + ReduceScatter(bwd)
        per_gpu = 2 * phi * factor

    # 假设 NVLink 600 GB/s 双向带宽
    bw = 600e9
    time_ms = (per_gpu / bw) * 1000

    return {
        "per_gpu_gb": per_gpu / 1e9,
        "total_gb": per_gpu * num_gpus / 1e9,
        "time_ms": time_ms,
        "time_percent": time_ms / 100,  # 假设总步时 ~100ms
    }

# 7B 模型, 8 卡
cost = communication_cost(14, 8, "ddp")
print(f"DDP: {cost['per_gpu_gb']:.1f} GB/卡, "
      f"通信占比 ~{cost['time_percent']:.0f}%")

# 7B 模型, 64 卡 — 通信占比爆炸
cost = communication_cost(14, 64, "ddp")
print(f"DDP 64卡: {cost['per_gpu_gb']:.1f} GB/卡, "
      f"通信占比 ~{cost['time_percent']:.0f}% ← 严重瓶颈")
卡数DDP 通信量/卡通信耗时 (NVLink)通信占比
4~21 GB~35 ms~35%
8~24.5 GB~41 ms~41%
16~26.3 GB~44 ms~44%
64~27.6 GB~46 ms~46%

关键结论:DDP 通信量随卡数增长趋于饱和(N−1N→1),但通信占比始终很高。这就是为什么大模型训练必须引入 TP/PP 来减少 DP 维度。


ZeRO — 优化器状态切分 ​

ZeRO 的三个阶段 ​

DeepSpeed ZeRO (Zero Redundancy Optimizer) 通过切分优化器状态、梯度和参数,逐步消除数据并行中的显存冗余。

mermaid
graph TD
    subgraph "ZeRO 三个阶段"
        BASE["基础数据并行<br/>每卡完整副本"] --> Z1["ZeRO-1<br/>切分优化器状态<br/>(β₁, β₂, momentum)"]
        Z1 --> Z2["ZeRO-2<br/>+ 切分梯度"]
        Z2 --> Z3["ZeRO-3<br/>+ 切分模型参数<br/>几乎线性扩展"]
    end

    style BASE fill:#f39c12,color:#fff
    style Z3 fill:#2ecc71,color:#fff

显存节省对比 ​

阶段优化器状态梯度参数单卡 7B 显存能训多大
无并行完整完整完整56 GB~10B
ZeRO-1切分完整完整42 GB~14B
ZeRO-2切分切分完整28 GB~20B
ZeRO-3切分切分切分14 GB~100B+

DeepSpeed ZeRO 配置 ​

json
{
    "train_batch_size": 64,
    "gradient_accumulation_steps": 4,
    "optimizer": {
        "type": "AdamW",
        "params": {
            "lr": 1e-4,
            "betas": [0.9, 0.95],
            "weight_decay": 0.1
        }
    },
    "zero_optimization": {
        "stage": 3,
        "offload_optimizer": {
            "device": "cpu",
            "pin_memory": true
        },
        "offload_param": {
            "device": "cpu",
            "pin_memory": true
        },
        "overlap_comm": true,
        "contiguous_gradients": true,
        "reduce_bucket_size": 5e7,
        "stage3_prefetch_bucket_size": 5e7,
        "stage3_param_persistence_threshold": 1e5,
        "sub_group_size": 1e9
    },
    "fp16": {
        "enabled": true
    }
}
bash
# DeepSpeed 启动命令
deepspeed --num_gpus=8 train.py \
  --deepspeed ds_config_zero3.json \
  --model_name Qwen/Qwen4-7B \
  --per_device_train_batch_size 8 \
  --gradient_accumulation_steps 4

模型并行 (Model Parallelism) ​

Tensor Parallelism (TP) — 层内切分 ​

将一个 Transformer 层内的权重矩阵沿列/行切分到多张 GPU。

mermaid
graph TD
    subgraph "Tensor Parallelism: 列切分 (Column-wise)"
        X["输入 x<br/>(B, L, d_model)"] --> G0["GPU 0: A[:, :d/2]<br/>y₀ = x · A₀"]
        X --> G1["GPU 1: A[:, d/2:]<br/>y₁ = x · A₁"]
        G0 --> CAT["拼接: y = [y₀ | y₁]"]
        G1 --> CAT
    end

    subgraph "Tensor Parallelism: 行切分 (Row-wise)"
        X2["输入 x"] --> SPLIT["切分: x₀, x₁"]
        SPLIT --> G0R["GPU 0: y₀ = x₀ · B₀"]
        SPLIT --> G1R["GPU 1: y₁ = x₁ · B₁"]
        G0R --> REDUCE["AllReduce: y = y₀ + y₁"]
        G1R --> REDUCE
    end

    style CAT fill:#3498db,color:#fff
    style REDUCE fill:#e74c3c,color:#fff
切分方式前向传播反向传播通信量
列切分输入复制,输出拼接梯度 AllReducefwd: AllGather, bwd: ReduceScatter
行切分输入切分,输出求和梯度切分fwd: AllReduce, bwd: AllGather
python
"""
Tensor Parallelism 的列切分实现
"""
import torch
import torch.nn as nn
import torch.distributed as dist

class ColumnParallelLinear(nn.Module):
    """列切分线性层"""

    def __init__(self, in_features, out_features, world_size, rank):
        super().__init__()
        self.rank = rank
        self.world_size = world_size
        assert out_features % world_size == 0

        local_out = out_features // world_size
        self.weight = nn.Parameter(torch.randn(local_out, in_features) * 0.02)

    def forward(self, x):
        # 每张卡计算: y_local = x @ W_local^T
        y_local = x @ self.weight.T  # (B, seq, local_out)

        # AllGather: 收集所有卡的输出并拼接
        y_list = [torch.zeros_like(y_local) for _ in range(self.world_size)]
        dist.all_gather(y_list, y_local)
        y = torch.cat(y_list, dim=-1)  # (B, seq, out_features)

        return y

Pipeline Parallelism (PP) — 层间切分 ​

将模型按层切分,每张 GPU 负责连续的若干层。流水线中,micro-batch 在不同 GPU 间流动。

mermaid
graph LR
    subgraph "Pipeline: 4 GPU, 4 micro-batches"
        MB1["MB1"] --> L0["GPU0<br/>层1-8"]
        L0 --> L1["GPU1<br/>层9-16"]
        L1 --> L2["GPU2<br/>层17-24"]
        L2 --> L3["GPU3<br/>层25-32"]
        L3 --> LOSS["Loss"]
    end

    NOTE["关键: micro-batch 实现流水线并行<br/>避免 GPU 空闲等待"]

3D 并行 — 组合使用 ​

mermaid
graph TD
    subgraph "3D 并行架构"
        NODE1["节点 1 (8 GPU)"] --> |"DP: 数据副本"| NODE2["节点 2 (8 GPU)"]

        subgraph "单节点内"
            G0["GPU 0"] --> |"TP: 层内切分"| G1["GPU 1"]
            G0 --> |"PP: 层间切分"| G4["GPU 4-7"]
            G4 --> G5["GPU 5"]
            G4 --> G6["GPU 6"]
            G4 --> G7["GPU 7"]
        end
    end

    style NODE1 fill:#3498db,color:#fff
    style G0 fill:#e74c3c,color:#fff
    style G4 fill:#2ecc71,color:#fff
并行策略拆分维度适用规模通信模式
DP数据 batch2-16 卡AllReduce 梯度
TP单层权重矩阵2-8 卡 (单机)AllReduce/AllGather 激活
PP模型层2-N 卡P2P 发送激活/梯度
ZeRO-3参数/梯度/优化器N 卡AllGather + ReduceScatter

生产环境典型配置:TP=2 + PP=4 + DP=16 = 128 GPU。先用 PP 和 TP 满足单模型副本的显存需求,再用 DP 加速。


FSDP — PyTorch 原生 ZeRO-3 ​

FSDP (Fully Sharded Data Parallel) 是 PyTorch 官方实现的 ZeRO-3 等效方案。

python
"""
PyTorch FSDP 完整示例
"""
import torch
import torch.nn as nn
from torch.distributed.fsdp import (
    FullyShardedDataParallel as FSDP,
    MixedPrecision,
    ShardingStrategy,
    BackwardPrefetch,
)
from torch.distributed.fsdp.wrap import (
    transformer_auto_wrap_policy,
    size_based_auto_wrap_policy,
)

def create_fsdp_model(model, rank):
    """包装模型为 FSDP"""

    # 混合精度配置
    mixed_precision = MixedPrecision(
        param_dtype=torch.bfloat16,   # 参数用 BF16
        reduce_dtype=torch.float32,   # 梯度归约用 FP32
        buffer_dtype=torch.float32,
    )

    # 自动包装策略: 按 Transformer 层包装
    auto_wrap_policy = transformer_auto_wrap_policy

    model = FSDP(
        model,
        auto_wrap_policy=auto_wrap_policy,
        mixed_precision=mixed_precision,
        sharding_strategy=ShardingStrategy.FULL_SHARD,
        backward_prefetch=BackwardPrefetch.BACKWARD_PRE,
        device_id=rank,
        cpu_offload=False,
    )

    return model


# 启动: torchrun --nproc_per_node=8 train_fsdp.py
def train_fsdp():
    import torch.distributed as dist
    dist.init_process_group(backend="nccl")

    local_rank = int(os.environ["LOCAL_RANK"])
    torch.cuda.set_device(local_rank)

    # 创建模型并 FSDP 包装
    from transformers import AutoModelForCausalLM
    model = AutoModelForCausalLM.from_pretrained(
        "Qwen/Qwen4-7B",
        torch_dtype=torch.bfloat16,
    )
    model = create_fsdp_model(model, local_rank)

    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)

    for epoch in range(3):
        for batch in dataloader:
            # FSDP 自动处理前向的参数 AllGather 和
            # 反向的梯度 ReduceScatter
            loss = model(**batch).loss
            loss.backward()
            optimizer.step()
            optimizer.zero_grad()

    dist.destroy_process_group()

通信原语速查 ​

原语操作用途
AllReduce所有卡求和/平均,结果分发到所有卡梯度同步
AllGather收集所有卡的数据,拼接后分发TP 列切分的输出拼接
ReduceScatter先求和再按卡切分分发ZeRO-3 梯度归约
Broadcast一张卡的数据复制到所有卡模型参数初始化
P2P (Send/Recv)点对点传输Pipeline Parallelism

分布式训练 Checklist ​

检查项说明
✅ NCCL 版本 ≥ 2.18nccl-tests 验证带宽
✅ NVLink/NVSwitch 可用nvidia-smi topo -m 查看拓扑
✅ 梯度累积对齐global_batch = per_device × grad_accum × num_gpus
✅ 学习率缩放lr 按 sqrt 或 linear 缩放,warmup 比例不变
✅ 随机种子同步所有 rank 使用相同 seed
✅ Checkpoint 一致性保存/加载时处理分片状态的聚合

参考 ​

批注模式

💬 文章评论

暂无评论,来说点什么吧 👇

编程学习笔记