如何在CPU上加载GPU训练的PyTorch模型?全模型加载为何报错?
PyTorch模型跨设备加载问题解答
1. 报错原因说明
不是因为保存了GPU专属数据,你遇到的报错核心原因是训练时使用了torch.nn.DataParallel(DP)对模型做了并行封装:
保存完整模型时会把整个DP实例的所有属性序列化存储,其中包含固定属性src_device_obj,默认值是训练时的设备cuda:0。map_location参数只能修改模型参数的存储设备,不会修改DP实例的内置属性,所以前向推理时DP的校验逻辑依然要求所有参数在cuda:0上,触发设备不匹配错误。
如果需要用加载完整模型的方式正常运行,可以加载后取出DP包裹的原始模型:
self.__device = torch.device('cpu') self.__model = torch.load(self.model_path, map_location=self.__device) # 提取DP封装的实际模型 if isinstance(self.__model, torch.nn.DataParallel): self.__model = self.__model.module
2. 加载方案选择建议
不推荐常规场景下使用加载完整模型的方案,优先选择加载state_dict的方案,原因如下:
- 兼容性更强:完整模型的序列化和训练时的模型类定义强绑定,加载环境的模型类路径、类结构(比如新增/删除层、修改参数名)只要和训练时有差异,就会加载失败;
state_dict只存储参数,只要模型结构匹配就能加载,灵活性更高。 - 存储成本更低:
state_dict仅存储可训练参数和缓冲层数据,文件体积远小于完整模型文件。 - 适配性更广:涉及迁移学习、模型微调、跨环境部署的场景下,
state_dict的适配成本远低于完整模型。
只有自用测试、完全不需要跨环境迁移的极小场景可以临时使用完整模型加载方案。
内容的提问来源于stack exchange,提问作者Bochjamin
相关产品推荐
相关产品推荐

