You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.29 04:45:03