关于大规模分布式训练的几点记录

上个月接手一个图像分类项目,训练数据大概 10TB,模型是 ResNet152 的魔改版。在实际项目中,我们遇到几个明显的限制:

  1. 显存限制:batch size 上不去,梯度不够稳定
  2. 时间限制:项目 deadline 紧,等不起单机慢吞吞
  3. 模型规模:后来尝试更大的模型,单卡根本装不下

简单计算一下:假设我们要训练一个需要 100GB 显存的模型,单张 3090 只有 24GB 显存,即使开启梯度累积,效率也会大幅下降。

为什么需要分布式训练

先说清楚一个问题:分布式训练不是万能药。

在实际项目中,我们遇到几个明显的限制:

  1. 显存限制:batch size 上不去,梯度不够稳定
  2. 时间限制:项目 deadline 紧,等不起单机慢吞吞
  3. 模型规模:后来尝试更大的模型,单卡根本装不下

简单计算一下:假设我们要训练一个需要 100GB 显存的模型,单张 3090 只有 24GB 显存,即使开启梯度累积,效率也会大幅下降。这时候多卡几乎是唯一选择。

单机多卡:第一步

单机多卡是最简单的分布式形式,PyTorch 自带的 DistributedDataParallel (DDP) 就能搞定。

环境准备

# 检查 GPU 状态
nvidia-smi

# 输出示例
# +-----------------------------------------------------------------------------+
# | NVIDIA-SMI 515.65.01    Driver Version: 515.65.01    CUDA Version: 12.0     |
# |-------------------------------+----------------------+----------------------+
# | GPU  Name        Persistence-M| Bus-Id        Disp.A | Volatile Uncorr. ECC |
# | Fan  Temp  Perf  Pwr:Usage/Cap|         Memory-Usage | GPU-Util  Compute M. |
# |===============================+======================+======================|
# |   0  NVIDIA GeForce ...  Off  | 00000000:01:00.0  On |                  N/A |
# | 30%   42C    P8    15W / 350W |   1234MiB / 24576MiB |      0%      Default |
# +-------------------------------+----------------------+----------------------+

基础代码

import torch
import torch.distributed as dist
import torch.multiprocessing as mp
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(
        backend='nccl',  # NVIDIA GPU 用 nccl,CPU 用 gloo
        rank=rank,
        world_size=world_size
    )

def cleanup():
    dist.destroy_process_group()

def train(rank, world_size):
    print(f"Running DDP on rank {rank}.")

    setup(rank, world_size)

    # 创建模型并移到当前 GPU
    model = Model().to(rank)
    ddp_model = DDP(model, device_ids=[rank])

    # 数据加载器需要使用 DistributedSampler
    dataset = YourDataset()
    sampler = torch.utils.data.distributed.DistributedSampler(
        dataset,
        num_replicas=world_size,
        rank=rank,
        shuffle=True
    )
    dataloader = torch.utils.data.DataLoader(
        dataset,
        batch_size=32,
        sampler=sampler
    )

    optimizer = torch.optim.Adam(ddp_model.parameters())

    for epoch in range(epochs):
        sampler.set_epoch(epoch)  # 确保每个 epoch 数据打乱不同

        for batch_idx, (data, target) in enumerate(dataloader):
            data, target = data.to(rank), target.to(rank)

            optimizer.zero_grad()
            output = ddp_model(data)
            loss = criterion(output, target)
            loss.backward()
            optimizer.step()

    cleanup()

if __name__ == "__main__":
    world_size = torch.cuda.device_count()
    mp.spawn(train, args=(world_size,), nprocs=world_size, join=True)

踩坑记录

坑一:NCCL 初始化失败

现象:RuntimeError: NCCL error in: /pytorch/torch/lib/c10d/ProcessGroupNCCL.cpp:1300, unhandled system error

原因:多块 GPU 之间的通信可能因为 PCIe 插槽、系统资源等问题失败。

解决:检查 GPU 的物理连接方式,最好插在同一个 PCIe 控制器下的插槽:

# 查看拓扑关系
nvidia-smi topo -m

# GPU0    GPU1    GPU2    GPU3    GPU4    GPU5    GPU6    GPU7    CPU Affinity    NUMA Affinity
# GPU0     X      PHB     PHB     PHB     PHB     PHB     PHB     PHB     0-15,32-47    0
# GPU1    PHB      X      PHB     PHB     PHB     PHB     PHB     PHB     0-15,32-47    0

# PHB 表示在同一 PCIe 主桥上,通信效率最高
# SOC 表示需要跨 PCIe 主桥,效率会低一些
# SYS 表示需要跨 CPU,效率最低

如果 GPU 之间通信效率低,考虑重新调整硬件连接或者接受性能损失。

坑二:Sampler 忘记设置 epoch

现象:每个 epoch 训练的数据完全一样,模型不收敛。

原因:DistributedSampler 默认使用相同的随机种子,不设置 epoch 会导致数据打乱方式相同。

解决:在每个 epoch 开始时调用 sampler.set_epoch(epoch)

for epoch in range(epochs):
    sampler.set_epoch(epoch)  # 这行很重要
    for batch_idx, (data, target) in enumerate(dataloader):
        # ...

坑三:多卡显存占用不均衡

现象:某些 GPU 显存占用很高,其他 GPU 显存占用很低。

原因:可能是数据分布不均或者某些 batch 的样本特别大。

解决:开启 gradient checkpointing 或调整数据加载策略:

# Gradient Checkpointing
from torch.utils.checkpoint import checkpoint

def forward_with_checkpointing(x):
    return checkpoint(self.custom_forward, x)

def custom_forward(self, x):
    return self.layers(x)

多机多卡:真正的分布式

单机多卡解决不了我们的问题,4 张 3090 仍然不够。于是决定上多机集群。

集群架构

我们的集群配置:

  • 4 台训练节点,每台 4 张 A100 (40GB)
  • 1 台管理节点,负责任务调度和监控
  • 100Gbps InfiniBand 网络(这个很关键)
  • NFS 共享存储,存储训练数据和模型
graph LR A[管理节点] --> B[训练节点1] A --> C[训练节点2] A --> D[训练节点3] A --> E[训练节点4] B <-->|InfiniBand| C C <-->|InfiniBand| D D <-->|InfiniBand| E B --> F[NFS共享存储] C --> F D --> F E --> F

网络配置

InfiniBand 对于分布式训练至关重要。我们试过普通的千兆以太网,训练效率惨不忍睹。

# 检查 InfiniBand 状态
ibstat

# 输出示例
# CA 'mlx5_0'
#         CA type: MT4125
#         Number of ports: 1
#         Port 1:
#                 State: Active
#                 Physical state: LinkUp
#                 Rate: 100
#                 Base lid: 2
#                 LMC: 0
#                 SM lid: 1
#                 Capability mask: 0x2658e848
#                 Port GUID: 0x506b4b03005b8c80

# 测试网络带宽
ib_write_bw -d mlx5_0 -i 1 -s 1073741824
# 理论上应该接近 100Gbps

启动脚本

多机环境下,启动方式需要调整:

#!/bin/bash

# 启动脚本:launch_cluster.sh

NODES=("node1" "node2" "node3" "node4")
NUM_GPUS_PER_NODE=4
MASTER_ADDR="${NODES[0]}"
MASTER_PORT=29500

# 启动所有节点
for NODE in "${NODES[@]}"; do
    ssh $NODE "python train.py \
        --nodes ${#NODES[@]} \
        --gpus $NUM_GPUS_PER_NODE \
        --node-rank $((i)) \
        --master-addr $MASTER_ADDR \
        --master-port $MASTER_PORT" &
    ((i++))
done

wait

对应的 Python 代码:

import argparse

parser = argparse.ArgumentParser()
parser.add_argument('--nodes', type=int, default=1)
parser.add_argument('--gpus', type=int, default=1)
parser.add_argument('--node-rank', type=int, default=0)
parser.add_argument('--master-addr', type=str, default='localhost')
parser.add_argument('--master-port', type=str, default='12355')
args = parser.parse_args()

def setup():
    dist.init_process_group(
        backend='nccl',
        init_method=f'tcp://{args.master_addr}:{args.master_port}',
        world_size=args.nodes * args.gpus,
        rank=args.node_rank * args.gpus + local_rank
    )

def main():
    local_rank = int(os.environ['LOCAL_RANK'])
    torch.cuda.set_device(local_rank)

    setup()

    # 其余训练逻辑和单机多卡相同
    # ...

if __name__ == '__main__':
    main()

踩坑记录

坑一:SSH 免密登录配置

现象:启动脚本执行后,其他节点没有反应。

原因:SSH 免密登录没有配置好,脚本无法在其他节点上执行。

解决:配置 SSH 免密登录:

# 在管理节点生成密钥对
ssh-keygen -t rsa -b 4096

# 将公钥复制到所有训练节点
for NODE in node1 node2 node3 node4; do
    ssh-copy-id user@$NODE
done

# 测试免密登录
ssh node1 "hostname"

坑二:NFS 性能瓶颈

现象:GPU 利用率不高,经常等待数据。

原因:NFS 共享存储的 I/O 性能不够,数据加载成为瓶颈。

解决:将数据缓存到本地 SSD:

class LocalCacheDataset(Dataset):
    def __init__(self, remote_path, cache_dir='/tmp/cache'):
        self.remote_path = remote_path
        self.cache_dir = cache_dir
        os.makedirs(cache_dir, exist_ok=True)
        self.cache = {}

    def __getitem__(self, idx):
        if idx in self.cache:
            return self.cache[idx]

        # 第一次从 NFS 加载
        data = self._load_from_remote(idx)

        # 缓存到本地
        cache_path = os.path.join(self.cache_dir, f'{idx}.pkl')
        with open(cache_path, 'wb') as f:
            pickle.dump(data, f)

        self.cache[idx] = data
        return data

坑三:节点间通信失败

现象:训练过程中某个节点突然失联,整个训练崩溃。

原因:网络抖动或者某个节点负载过高导致超时。

解决:增加超时时间和重试机制:

# 增加初始化超时时间
dist.init_process_group(
    backend='nccl',
    timeout=timedelta(minutes=10)  # 默认是30分钟,可以根据实际情况调整
)

# 监控节点健康状态
def monitor_node_health():
    while True:
        # 检查 NCCL 通信是否正常
        try:
            dist.all_reduce(torch.tensor([0.0]), op=dist.ReduceOp.SUM)
        except Exception as e:
            print(f"节点通信异常: {e}")
            # 可以选择重启训练或者记录日志

        time.sleep(60)

性能优化技巧

混合精度训练

from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()

for data, target in dataloader:
    optimizer.zero_grad()

    with autocast():
        output = model(data)
        loss = criterion(output, target)

    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

混合精度训练可以:

  • 减少显存占用,可以增大 batch size
  • 利用 Tensor Core 加速计算
  • 通常可以带来 2-3 倍的性能提升

梯度累积

accumulation_steps = 4

for i, (data, target) in enumerate(dataloader):
    with autocast():
        output = model(data)
        loss = criterion(output, target) / accumulation_steps

    scaler.scale(loss).backward()

    if (i + 1) % accumulation_steps == 0:
        scaler.step(optimizer)
        scaler.update()
        optimizer.zero_grad()

梯度累积可以模拟更大的 batch size,在不增加显存的情况下获得大 batch size 的训练效果。

数据预取和多线程加载

dataloader = torch.utils.data.DataLoader(
    dataset,
    batch_size=32,
    num_workers=8,  # 多线程预取
    pin_memory=True,  # 锁页内存,加速 GPU 传输
    persistent_workers=True  # 保持 worker 进程,避免重复创建
)

性能监控

分布式训练的监控比单机复杂很多,需要关注多节点的状态。

使用 NVIDIA Nsight

# 安装 Nsight Systems
pip install nvidia-nsight-systems

# 运行训练并收集性能数据
nsys profile --stats=true --output=profile_report python train.py

# 分析报告
nsys stats profile_report.nsys-rep

自定义监控指标

import time
import torch.distributed as dist

class TrainingMonitor:
    def __init__(self):
        self.step_times = []
        self.data_times = []
        self.compute_times = []

    def record_step(self, data_time, compute_time):
        self.step_times.append(data_time + compute_time)
        self.data_times.append(data_time)
        self.compute_times.append(compute_time)

    def report(self):
        # 只在 rank 0 上输出
        if dist.get_rank() == 0:
            avg_step = sum(self.step_times) / len(self.step_times)
            avg_data = sum(self.data_times) / len(self.data_times)
            avg_compute = sum(self.compute_times) / len(self.compute_times)

            print(f"平均步时间: {avg_step:.2f}s")
            print(f"数据加载时间: {avg_data:.2f}s ({avg_data/avg_step:.1%})")
            print(f"计算时间: {avg_compute:.2f}s ({avg_compute/avg_step:.1%})")

            # 数据加载占比过高说明是 I/O 瓶颈
            # 计算时间占比过高说明是计算瓶颈

实际效果对比

经过一系列优化,我们的训练效率有了明显提升:

配置单个 epoch 时间加速比
单张 309012 小时1.0x
4 张 3090(单机)3.5 小时3.4x
16 张 A100(4 机)45 分钟16.0x

把三种配置的 epoch 耗时和加速比放在一张图里,比单看表格更直观——16 卡集群几乎追上了线性加速的理论上限。

ResNet152 分布式训练:单卡 3090、4 卡 3090 与 16 卡 A100 集群的 epoch 耗时和加速比对比

这说明通信、I/O 和负载均衡的优化基本到位,硬件扩展带来的收益没有被工程开销大量吃掉。

虽然理论加速比是线性的,但实际中会有各种开销:

  • 梯度同步通信开销
  • 数据加载瓶颈
  • 负载不均衡
  • 网络延迟

我们的 16 张 A100 基本达到了线性加速,说明优化比较到位。

一些经验总结

  1. 先优化单机再上集群:单机上的性能问题会在集群上被放大
  2. 网络很关键:InfiniBand 不是必须的,但普通以太网会严重限制性能
  3. 数据存储要考虑 I/O:NFS 共享存储可能成为瓶颈,本地缓存是有效的解决方案
  4. 监控比调优更重要:没有监控的调优是盲人摸象
  5. 容错机制必不可少:集群环境下节点故障是常态,需要有恢复机制

最后说一句

分布式训练看着复杂,但核心思想很简单:把大任务拆成小任务,并行执行。真正麻烦的是工程实践中的各种细节问题。

从单机到集群的迁移,不只是改几行代码,更是整个基础设施的升级。硬件、网络、存储、监控,每个环节都不能忽视。

但好处也是实实在在的:原本一周的训练现在半天就能完成,给调参和实验留出了充足时间。项目最终提前一周交付,客户也很满意。

技术就是这样,投入成本解决问题,用效率提升证明价值。只是别忘了,分布式训练不是银弹,有些问题可能需要从算法或模型层面解决,而不是堆硬件。


这次分布式训练实践花了一个月,从最初的单机训练到最终的 16 卡集群,中间踩过不少坑。现在回过头看,很多问题都有迹可循,只是当时没有足够的经验。希望这篇记录能帮到有类似需求的同学。

版权声明: 本文首发于 指尖魔法屋-关于大规模分布式训练的几点记录https://blog.thinkmoon.cn/post/186-large-scale-distributed-training-single-to-cluster/) 转载或引用必须申明原指尖魔法屋来源及源地址!