如何将PyTorch TPU训练的模型保存为可在CPU环境加载的权重
跨环境加载TPU训练PyTorch模型权重的解决方案
报错原因
你直接在TPU环境下调用torch.save保存模型状态字典时,字典内的所有张量都绑定了XLA后端,没有安装pytorch_xla依赖的环境无法识别该后端类型,因此触发aten::empty_strided相关报错。
方案1:保存时预处理(最推荐)
在TPU训练环境执行保存操作前,先将所有模型参数转移到CPU上,剥离XLA后端绑定,再保存权重。这样输出的权重文件天然兼容所有PyTorch支持的环境:
# TPU训练环境下的保存代码 cpu_state_dict = {k: v.cpu() for k, v in model.state_dict().items()} file_name = "model_params" torch.save(cpu_state_dict, file_name)
保存完成后,任意环境都可以用常规加载代码直接加载,无需额外调整。
方案2:已保存XLA权重的加载补救
如果已经保存了绑定XLA后端的权重,无需重新训练导出,加载时指定map_location参数强制将张量映射到目标设备(CPU/CUDA均可),绕过XLA后端解析即可:
# 无TPU环境下的加载代码 model.load_state_dict(torch.load(file_name, map_location=torch.device('cpu'))) # 若要加载到GPU使用,也可以直接指定为 map_location='cuda'
注:该方案要求PyTorch版本≥1.12,低版本可能存在兼容问题。
内容的提问来源于stack exchange,提问作者Yedidya kfir
相关产品推荐
相关产品推荐

