Skip to content

PyTorch 分布式计算常见问题与解决方案

内容整理自学习笔记,仅供面试备考参考;不构成录用、培训或考试承诺。

1. DistributedDataParallel 显存分布不均衡问题

1.1 问题描述

使用 DistributedDataParallel 时,每个进程独立运行在一个 GPU 上,理论上各卡显存应该均匀分布。但有时其他进程会在 0 号卡上额外占用显存,导致 0 号卡显存瓶颈,引发 CUDA Out-of-Memory 错误。

# 期望:均匀分布
GPU0: [████████]  GPU1: [████████]  GPU2: [████████]  GPU3: [████████]

# 实际问题:0 号卡额外占用
GPU0: [████████████]  GPU1: [████████]  GPU2: [████████]  GPU3: [████████]

1.2 问题定位

问题根源在于 torch.load() 默认将数据加载到 0 号卡:

python
# 错误写法:数据默认 load 到 cuda:0
checkpoint = torch.load("checkpoint.pth")
model.load_state_dict(checkpoint["state_dict"])

注意: 在 Distributed 模式下,代码分别运行在多个 GPU 上,但 torch.load() 默认将数据放到 0 号卡,导致所有进程都在 0 号卡占用一部分显存。

1.3 解决方法

将 load 的数据映射到 CPU:

python
# 正确写法:映射到 CPU
checkpoint = torch.load("checkpoint.pth", map_location=torch.device('cpu'))
model.load_state_dict(checkpoint["state_dict"])

2. 自研数据接口导致程序卡死

2.1 问题描述

使用 PyTorch 实现同步梯度更新时,如果数据接口是自研的,可能出现第一个 epoch 结束时程序卡死在 torch.all_reduce() 上,且没有任何报错信息。

2.2 根因分析与解决

根因: 每张卡分配的 batch 数量不一致,某张卡少了一个 batch,其他卡在 AllReduce 时一直等待,导致程序卡死。

关键原则: 必须保证每张卡分配的 batch 数量完全一致。

3. 微调大模型时多卡卡死问题

3.1 问题描述

微调大模型时,单机 2 卡训练正常,但使用 4 卡及以上时,程序卡在读完数据和开始训练之间,无任何输出。

3.2 排查步骤

  1. 确认 GPU 通信正常: 检查所有 GPU 是否能正常使用和通信
  2. 检查 batch 分配: 是否是 batch size 分配问题导致某张卡无限等待
  3. 最小化测试: 只留 4 条数据,每张卡只跑一条,验证流程是否正常

4. 常见坑点总结

问题症状根因解决方法
显存不均衡0 号卡 OOMtorch.load() 默认加载到 cuda:0使用 map_location='cpu'
epoch 结尾卡死all_reduce 无限等待各卡 batch 数量不一致确保 batch 分配均匀
多卡训练卡死读完数据后无响应通信或 batch 分配问题最小化数据测试+检查通信