PyTorch分布式训练CUDA显存不足,显存占用数值总和不符问题
问题根因分析
你遇到的OOM报错属于典型的非PyTorch主进程占用显存导致的异常,报错里PyTorch统计的仅为当前进程自身的显存占用,不会统计同GPU上其他进程、上下文占用的显存,可能的原因如下:
- 同GPU存在其他无关进程占用显存:包括本机其他用户的训练任务、你之前异常退出未杀死的残留进程、桌面环境GPU加速、监控进程等,占走了15G左右的显存,仅剩100多M给当前训练进程使用。
- 分布式启动参数配置错误:使用Torch Distributed Elastic时如果
--nproc_per_node设置大于单卡数量,或者未正确设置GPU可见性,会导致多个训练进程同时绑定到GPU 0上,每个进程的PyTorch仅统计自身的显存占用,累加后占满整卡。 - DataLoader worker进程显存泄漏:你设置了2个DataLoader worker,如果在数据集
__getitem__逻辑中存在将张量移动到GPU的操作,worker进程会持有这部分显存且不会被主进程的显存统计覆盖,多个worker累积后占满显存。 - 残留CUDA上下文占用显存:之前的训练任务异常退出时未正确销毁CUDA上下文,这部分显存会被系统锁定无法分配给新进程,也不会被新的PyTorch进程统计到。
- 容器资源限制:如果在Docker、K8s等容器环境运行,容器的GPU显存配额被设置为远小于物理卡容量,PyTorch仅能使用配额内的显存,就会出现物理卡显存15G但实际可用不足的情况。
可行调试方案
- 首先执行
nvidia-smi命令查看目标GPU的全局显存占用和进程列表,确认是否有无关进程占用显存,找到对应PID后执行kill -9 <PID>清理即可。 - 验证分布式启动配置,确认启动命令中
--nproc_per_node参数设置为2,和单机器的GPU数量匹配,同时可以在代码开头加上import os; print(os.environ['CUDA_VISIBLE_DEVICES']),确认每个进程的可见GPU是正确的。 - 排查DataLoader实现逻辑,确保所有数据加载、预处理操作都仅在CPU上执行,将张量移动到GPU的操作放到主进程的训练循环中;同时可以临时将
num_workers设为0验证是否是worker进程导致的问题。 - 如果
nvidia-smi没有显示占用显存的进程但显存还是被占满,执行lsof /dev/nvidia0查看所有占用GPU 0的进程,清理残留进程后仍无效的话,可以执行sudo nvidia-smi -r -i 0重置对应GPU,或者直接重启机器清理残留CUDA上下文。 - 容器环境下检查资源配置:Docker运行时确认没有设置错误的显存限制参数,
--shm-size至少设置为16G避免共享内存不足导致的异常显存占用;K8s环境下检查容器的nvidia.com/gpu资源申请和限制是否正确。 - 临时将batch size调整为1,运行最小测试用例,同时在代码中调用
print(torch.cuda.memory_summary())打印完整的显存统计信息,定位是否存在未释放的张量占用。
内容的提问来源于stack exchange,提问作者Martyna
相关产品推荐
相关产品推荐

