外观
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 排查步骤
- 确认 GPU 通信正常: 检查所有 GPU 是否能正常使用和通信
- 检查 batch 分配: 是否是 batch size 分配问题导致某张卡无限等待
- 最小化测试: 只留 4 条数据,每张卡只跑一条,验证流程是否正常
4. 常见坑点总结
| 问题 | 症状 | 根因 | 解决方法 |
|---|---|---|---|
| 显存不均衡 | 0 号卡 OOM | torch.load() 默认加载到 cuda:0 | 使用 map_location='cpu' |
| epoch 结尾卡死 | all_reduce 无限等待 | 各卡 batch 数量不一致 | 确保 batch 分配均匀 |
| 多卡训练卡死 | 读完数据后无响应 | 通信或 batch 分配问题 | 最小化数据测试+检查通信 |