AI分布式训练踩坑记录

显存占用一路狂奔到 23.5GB,然后崩了。

为什么需要分布式训练

问题很直接:模型太大,显存不够。

当前的情况:

  • 模型参数量:7B,fp16 权重约 14GB
  • 训练时需要的显存:模型权重 + 梯度 + 优化器状态 + 激活值
  • 单卡 3090 显存:24GB,训练时还会被 CUDA 上下文、碎片化吃掉一部分

算一下:权重 14GB + 梯度 14GB + 优化器状态 42GB(AdamW)= 70GB,这还没算激活值。一张卡根本不够。

显存占用大概是这样分布的:

7B 模型单卡训练显存占用对比

左边是单张卡的占用量,右边是 4 卡分摊后的情况——即便分摊了,激活值和优化器状态还是很吃力。

所以必须用多卡,甚至多机。

第一步:数据并行

先从最简单的开始——数据并行(Data Parallel)。

思路很直观:把模型复制到每张卡上,每张卡吃不同批次的数据,算完梯度后同步。

数据并行工作流程

每个 GPU 有自己的模型副本和数据批次,前向传播后梯度需要同步到所有卡上。

代码层面:

import torch
import torch.nn as nn
from torch.nn.parallel import DistributedDataParallel as DDP

def setup(rank, world_size):
    os.environ['MASTER_ADDR'] = 'localhost'
    os.environ['MASTER_PORT'] = '12355'
    dist.init_process_group("nccl", rank=rank, world_size=world_size)

def cleanup():
    dist.destroy_process_group()

def train(rank, world_size):
    setup(rank, world_size)

    model = Model().to(rank)
    ddp_model = DDP(model, device_ids=[rank])

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

    for data, label in dataloader:
        optimizer.zero_grad()
        output = ddp_model(data)
        loss = criterion(output, label)
        loss.backward()
        optimizer.step()

    cleanup()

数据并行的限制

  • 每张卡都要存完整模型副本,模型大小受单卡显存限制
  • 如果模型已经塞不满一张卡,数据并行就够用了
  • 但如果单张卡连模型都放不下,就得换方案

这次的数据并行测试确实快了:4 张卡相比 1 张卡,理论加速比是 4 倍,实际测下来大概 3.6 倍。 losses 0.1 左右的损耗主要是通信和同步开销。

但问题是:换更大的模型时,单卡还是不够。

第二步:模型并行

模型放不下单卡,就拆开。

模型并行(Model Parallel)的核心是:把模型的不同层(或者层内的不同参数)放在不同的设备上,前向传播时数据在这些设备间流动。

层间并行

最简单的方式是按层切分:

flowchart LR subgraph GPU0 A[Embedding] --> B[Transformer Layer 0-3] end subgraph GPU1 C[Transformer Layer 4-7] --> D[Transformer Layer 8-11] end subgraph GPU2 E[Transformer Layer 12-15] --> F[Transformer Layer 16-19] end subgraph GPU3 G[Transformer Layer 20-23] --> H[LM Head] end B --> C D --> E F --> G

代码层面:

class SplitModel(nn.Module):
    def __init__(self):
        super().__init__()
        self.part0 = nn.Sequential(...).to('cuda:0')
        self.part1 = nn.Sequential(...).to('cuda:1')
        self.part2 = nn.Sequential(...).to('cuda:2')
        self.part3 = nn.Sequential(...).to('cuda:3')

    def forward(self, x):
        x = self.part0(x.to('cuda:0'))
        x = self.part1(x.to('cuda:1'))
        x = self.part2(x.to('cuda:2'))
        x = self.part3(x.to('cuda:3'))
        return x

层间并行的坑

  • 层间通信开销大,每层都要跨设备传数据
  • 计算效率低,大部分时候只有一张卡在工作,其他卡在等
  • 负载不均衡,不同层的计算量可能差异很大

实际测下来,层间并行比单卡还慢,因为通信吃掉了所有加速。

张量并行

更高效的方式是张量并行(Tensor Parallel),把层内的参数矩阵切分。

以一个线性层 y = xW 为例:

# 原始矩阵 W shape: [hidden_size, hidden_size]

# 在两张卡上水平切分 W
W0 = W[:, :hidden_size//2]  # 在 GPU0
W1 = W[:, hidden_size//2:]  # 在 GPU1

# 输入 x 也需要对应切分
x0 = x[:, :hidden_size//2]
x1 = x[:, hidden_size//2:]

# 每张卡计算部分结果
y0 = x0 @ W0
y1 = x1 @ W1

# 最后 all-reduce 合并结果
y = y0 + y1

张量并行的优势

  • 同一层内的并行,计算更均衡
  • 通信只在层结束时发生,开销相对较小
  • 可以和数据并行叠加使用

这次用张量并行后,总算把 13B 模型在 8 张 3090 上跑起来了,每个 GPU 显存占用约 18GB,留了 6GB 给激活值。

第三步:多机集群

单机 8 张卡也到上限了,要想更大,就得搞多机。

网络准备

多机的关键在网络:带宽和延迟都很重要。

实践中的配置:

  • 网卡:Mellanox ConnectX-6, 200Gbps
  • 交换机:同样是 Mellanox,确保全线速
  • 网络拓扑:最好用 Fat-Tree,避免单点瓶颈

普通千兆网别想了,分布式训练的网络通信量是 PB 级的,千兆会被打爆。

环境配置

多机比单机复杂得多,需要:

# 1. 在所有节点安装相同的 PyTorch 版本
pip install torch==2.1.0 --index-url https://download.pytorch.org/whl/cu121

# 2. 确保所有节点能无密码 SSH 互相访问
ssh-keygen -t rsa
ssh-copy-id user@node1
ssh-copy-id user@node2

# 3. 同步代码和数据
rsync -avz /data/model user@node1:/data/
rsync -avz /data/model user@node2:/data/

# 4. 配置环境变量
export MASTER_ADDR="node1"  # 主节点 IP
export MASTER_PORT="29500"
export WORLD_SIZE=8  # 总进程数(2 节点 x 4 卡)
export NCCL_DEBUG=INFO  # 方便调试

启动脚本

#!/bin/bash
# launch_cluster.sh

NNODES=2
NODE_RANK=0
GPUS_PER_NODE=4

# 在每个节点上运行
torchrun \
    --nproc_per_node=$GPUS_PER_NODE \
    --nnodes=$NNODES \
    --node_rank=$NODE_RANK \
    --master_addr=$MASTER_ADDR \
    --master_port=$MASTER_PORT \
    train.py \
    --config config.yaml

多机训练的坑

  1. NCCL 通信问题

    • 报错:NCCL error: unhandled system error
    • 原因:网卡驱动版本不一致,或者 NCCL 环境变量没配对
    • 解决:更新所有节点到相同版本的驱动,设置 NCCL_IB_DISABLE=0
  2. 数据同步慢

    • 多机训练时,梯度同步跨节点,网络成了瓶颈
    • 解决:用梯度累积减少通信频率,或者用 ZeRO 优化器
  3. 节点故障恢复

    • 训练 3 天了,突然一个节点挂了
    • 解决:必须 checkpoint,定期保存模型和优化器状态

第四步:ZeRO 优化器

多机训练最耗时的就是通信。DeepSpeed 的 ZeRO(Zero Redundancy Optimizer)就是来解决这个问题的。

ZeRO 的核心思想:不要在每个进程上存一份完整的状态,而是切分它。

ZeRO-1

只切分优化器状态,每个进程只存 1/N 的状态,然后 all-gather 合并。

原始:每个进程存完整状态(权重 + 梯度 + 优化器状态)
ZeRO-1:每个进程只存 1/N 的优化器状态
节省:4 倍显存(对于 AdamW)

ZeRO-2

在 ZeRO-1 基础上,再切分梯度。

节省:8 倍显存
通信:每个 step 需要 all-gather + reduce-scatter

ZeRO-3

最激进:连模型权重都切分。

节省:N 倍显存(N 是进程数)
通信:每层前向传播都要 all-gather 权重
适用:超大模型,其他方法都塞不下时

ZeRO 各阶段显存占用对比

从基线到 ZeRO-3,每个 GPU 需要的状态逐渐减少,代价是通信开销增加。

实际配置:

import deepspeed

ds_config = {
    "train_batch_size": 32,
    "train_micro_batch_size_per_gpu": 2,
    "gradient_accumulation_steps": 4,
    "optimizer": {
        "type": "AdamW",
        "params": {
            "lr": 1e-4,
            "betas": [0.9, 0.999],
            "eps": 1e-8
        }
    },
    "fp16": {
        "enabled": True,
        "loss_scale": 0,
        "loss_scale_window": 1000,
        "initial_scale_power": 16,
        "hysteresis": 2,
        "min_loss_scale": 1
    },
    "zero_optimization": {
        "stage": 2,  # ZeRO-2
        "offload_optimizer": {
            "device": "cpu",  # 把优化器状态卸载到 CPU
            "pin_memory": True
        },
        "offload_param": {
            "device": "cpu"
        }
    },
    "gradient_clipping": 1.0,
    "prescale_gradients": False,
    "wall_clock_breakdown": False
}

model_engine, optimizer, _, _ = deepspeed.initialize(
    model=model,
    model_parameters=model.parameters(),
    config=ds_config
)

用 ZeRO-2 + CPU offload 后,同样的硬件能塞下 30B 模型,代价是每个 step 稍微慢一点,因为 CPU-GPU 通信。

实际结果

折腾了一圈,最终的效果:

配置模型大小总显存需求实测吞吐备注
单卡 30901.3B20GB15 samples/s基线
单机 4 卡数据并行1.3B80GB50 samples/s3.3x 加速
单机 4 卡张量并行7B72GB12 samples/s模型放大了
多机 8 卡 ZeRO-213B144GB8 samples/s通信开销明显
多机 8 卡 ZeRO-330B160GB3 samples/s能跑,但慢

几个关键观察:

  1. 数据并行在模型能塞下单卡时最划算,效率高,实现简单
  2. 张量并行适合中等大小的模型,比如 7B-13B
  3. ZeRO-3 能塞超大模型,但代价是吞吐量下降明显
  4. 多机训练的通信开销不可忽视,除非必要,先吃满单机

踩坑记录

过程中遇到的问题,按频率排序:

  1. CUDA OOM

    • 现象:训练一段时间后突然崩掉
    • 原因:显存碎片化,或者某个 batch 太大
    • 解决:用 torch.cuda.empty_cache() 手动清理,或者调小 batch size
  2. NCCL 超时

    • 现象:多机训练卡住,日志里看不到更新
    • 原因:某个节点的网络有问题,或者 NCCL 版本不匹配
    • 解决:换用稳定网络,确保所有节点环境一致
  3. 梯度爆炸/消失

    • 现象:loss 突然变成 NaN
    • 原因:学习率太大,或者模型初始化有问题
    • 解决:梯度裁剪,调小学习率,检查初始化
  4. Checkpoint 不兼容

    • 现象:恢复训练时报错
    • 原因:DeepSpeed 版本变了,或者配置变了
    • 解决:定期保存全量 checkpoint,不要只存最新的
  5. 数据加载慢

    • 现象:GPU 利用率上不去
    • 原因:CPU 处理数据跟不上 GPU
    • 解决:增加 worker 数量,用 faster 数据格式

结语

分布式训练不是银弹,它是权衡。

能单机解决的,就不要上多机。能用数据并行的,就不要搞模型并行。能用 ZeRO-2 的,就不要上 ZeRO-3。

因为每多一层复杂度,就多一层出错的可能,也多一层调试的成本。

但模型的大小还在往上走,硬件的瓶颈也在往前推。分布式训练终究是要搞的,只是搞之前要想清楚:你的瓶颈在哪里?是显存、计算、还是通信?

搞清楚了,方案自然就出来了。

折腾到现在,总算把一条路跑通了,但我知道这条路还会继续延伸,因为模型会更大,硬件会变,问题也会变。

这也是为什么值得写下来——至少下次遇到类似问题时,不用从零开始。

版权声明: 本文首发于 指尖魔法屋-AI分布式训练踩坑记录https://blog.thinkmoon.cn/post/272-ai-distributed-training-single-cluster-practice/) 转载或引用必须申明原指尖魔法屋来源及源地址!