加载PyTorch checkpoint遇CUDA版本错误,仅CPU可用如何解决?
解决CPU环境加载GPU训练的PyTorch检查点报错问题
这个报错的核心原因很明确:你之前在GPU环境下保存的检查点,所有模型参数张量都是绑定在GPU设备上的,而当前环境不仅没有可用GPU,CUDA驱动版本还和原环境不匹配,导致PyTorch尝试用GPU加载时失败。
解决方法其实很简单,只需要在torch.load()里加上map_location参数,指定把所有张量映射到CPU上:
# 替换原来的加载代码 checkpoint = torch.load(pathname, map_location=torch.device('cpu'))
额外注意事项
- 加载完检查点后,如果要使用模型进行推理,记得把模型也切换到CPU模式:
# 假设你已经定义了模型类 model = YourModelArchitecture() # 先把模型移到CPU model = model.to('cpu') # 再加载参数 model.load_state_dict(checkpoint['model_state_dict']) - 如果你的检查点里还保存了优化器的状态(比如
checkpoint['optimizer_state_dict']),这些状态也会被自动映射到CPU上。不过在纯CPU环境下,一般不需要继续训练,所以可以忽略这部分;如果真要后续训练,确保优化器也是在CPU上初始化的。
这个方法能完美绕过CUDA版本不兼容的问题,同时让你在CPU环境下顺利加载GPU训练的模型权重。
内容的提问来源于stack exchange,提问作者Tom Hale
相关产品推荐
相关产品推荐

