如何解决PyTorch DataParallel训练模型单GPU推理的CUDA无效设备序号错误
解决PyTorch DataParallel训练模型单GPU推理的CUDA设备序号错误问题
核心原因
用DataParallel训练的模型,参数会被存储在module.xxx的命名空间下,且部分张量绑定了多GPU训练时的设备索引。单GPU推理时若直接加载模型,batch大小大于1时会触发多设备分发逻辑,尝试访问不存在的GPU设备,从而抛出设备序号无效错误;而batch=1时因张量分发逻辑简化,未触发该问题。
具体解决步骤
剥离模型的module层
加载模型时,需移除DataParallel添加的module前缀:# 正确加载流程 model = YourModel() # 加载训练好的权重文件 state_dict = torch.load('trained_model.pth') # 移除参数名中的module.前缀 cleaned_state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()} model.load_state_dict(cleaned_state_dict) # 将模型移至单GPU model = model.to('cuda:0')若直接加载整个模型对象,可通过以下方式剥离module层:
model = torch.load('trained_model.pth') if isinstance(model, torch.nn.DataParallel): model = model.module model.to('cuda:0')强制映射权重到单GPU
加载权重时指定map_location,强制将所有张量映射到目标GPU,避免设备索引残留:state_dict = torch.load('trained_model.pth', map_location='cuda:0') cleaned_state_dict = {k.replace('module.', ''): v for k, v in state_dict.items()} model.load_state_dict(cleaned_state_dict) model.to('cuda:0')确保输入张量设备匹配
推理前必须将输入数据移至同一GPU,避免设备不匹配问题:# 假设inputs是你的批量输入数据 inputs = inputs.to('cuda:0') with torch.no_grad(): outputs = model(inputs)
验证方式
修改完成后设置batch size大于1进行推理,若不再出现CUDA: invalid device ordinal错误,则问题解决。
内容的提问来源于stack exchange,提问作者517453088qqcom
相关产品推荐
相关产品推荐

