Attempting to deserialize object on a CUDA device but torch.cuda.is_available() is False
·
当使用PyTorch在特定CUDA设备(如cuda:5)上训练并保存模型后,若在无相同GPU设备的环境中加载,会因设备元数据不匹配触发RuntimeError: Attempting to deserialize object on a CUDA device but torch.cuda.is_available() is False。根本原因是PyTorch序列化时隐式嵌入了原始设备信息,加载时默认尝试恢复至该设备。解决方法分为保存阶段预防和加载阶段适配两类,需根据实际场景选择:
一、保存阶段预防:强制模型移至CPU再保存
若提前规划跨设备部署,应在保存模型前主动将模型/权重转移至CPU,避免设备信息绑定。
1. 仅保存模型参数(推荐)
# 训练完成后,先将模型移至CPU再保存state_dict
model.cpu() # 确保模型在CPU上
torch.save(model.state_dict(), "model_weights.pth")
- 优点:生成的文件完全脱离设备依赖,可在任意环境直接加载。
- 加载方式(无需额外处理设备映射):
2. 保存完整检查点时同步处理
若需保存优化器状态等信息,同样需先移至CPU:
torch.save({
'model_state_dict': model.cpu().state_dict(), # 关键:强制CPU化
'optimizer_state_dict': optimizer.state_dict(),
'epoch': epoch,
}, "checkpoint.pth")
- 注意:
model.cpu()会修改原模型设备,若需继续训练,保存后应调用model.to(original_device)恢复。
二、加载阶段适配:动态指定目标设备
若模型已按原始设备保存(如cuda:5),可通过map_location强制重定向设备。
1. 通用加载方案(自动适配可用设备)
# 动态检测目标设备
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# 加载时重定向至目标设备
state_dict = torch.load("model_gpu5.pth", map_location=device)
model.load_state_dict(state_dict)
model.to(device).eval() # 确保模型整体处于目标设备
- 关键点:
map_location仅影响加载过程中的张量位置,需额外调用model.to(device)同步模型缓冲区(如BatchNorm的统计量)。- 必须显式指定
**map_location**,否则PyTorch会尝试恢复至原始设备(如cuda:5)。
2. 特殊场景处理
(1) 多GPU训练模型(含module.前缀)
若模型通过nn.DataParallel训练,state_dict键名含module.前缀,需清洗前缀:
state_dict = torch.load("multi_gpu_model.pth", map_location=device)
# 清洗module.前缀
state_dict = {k.replace("module.", ""): v for k, v in state_dict.items()}
model.load_state_dict(state_dict)
- 原因:单卡模型结构不含
module.前缀,直接加载会触发键名不匹配错误。
(2) Apple M系列芯片(MPS加速)
if torch.backends.mps.is_available():
device = torch.device("mps")
else:
device = torch.device("cpu")
state_dict = torch.load("model_gpu5.pth", map_location=device)
三、关键注意事项
**map_location**与**model.to()**需配合使用map_location仅控制加载时的张量位置,而模型缓冲区(如running_mean)可能仍保留在原始设备。必须调用**model.to(device)**确保整体一致性。- 避免保存完整模型对象
直接torch.save(model)会序列化整个类结构,导致跨环境兼容性问题。始终优先使用**state_dict**方式保存,提升可移植性与安全性。 - 生产环境最佳实践
- 保存时统一转至CPU:
torch.save(model.cpu().state_dict(), ...)。 - 加载时强制指定设备映射:
torch.load(..., map_location=device)。 - 记录元数据:在检查点中加入
device、pytorch_version等信息,便于调试。
- 保存时统一转至CPU:
总结:该错误本质是PyTorch序列化对设备拓扑的强依赖。预防优于修复——训练完成后立即将模型移至CPU再保存,可彻底避免设备不匹配问题;若已存在设备绑定模型,则必须通过map_location动态重定向设备,并清洗多GPU训练残留的module.前缀。核心原则是解耦模型权重与硬件环境。
更多推荐

所有评论(0)