Apple M2的MPS设备加载预训练权重遇float64不支持错误求解
解决MPS设备加载预训练权重时float64不支持的问题
你的思路完全可行——先将预训练权重加载到CPU转为float32,再迁移到MPS设备,这是解决该报错的标准方案。以下是具体实现方法:
方法一:加载权重时单独转换float64参数
这种方式只转换权重字典中属于float64的参数,保留其他类型参数的原始格式:
import torch # 1. 把预训练权重加载到CPU state_dict = torch.load("pretrained_weights.pth", map_location="cpu") # 2. 遍历权重字典,将所有float64类型的参数转为float32 for param_name in state_dict: if state_dict[param_name].dtype == torch.float64: state_dict[param_name] = state_dict[param_name].float() # 3. 加载权重到模型,再迁移到MPS model = YourModelClass() model.load_state_dict(state_dict) model.to("mps")
方法二:加载模型后整体转换为float32
如果模型所有参数都可以转为float32(大部分预训练模型都支持),可以直接对模型整体做类型转换:
import torch # 1. 初始化模型并加载权重到CPU model = YourModelClass() model.load_state_dict(torch.load("pretrained_weights.pth", map_location="cpu")) # 2. 将模型所有参数转为float32,再迁移到MPS model = model.float() model.to("mps")
额外注意事项
- 运行模型时,输入的张量也必须是float32类型并迁移到MPS,否则仍会触发类型不兼容报错:
input_tensor = input_tensor.float().to("mps") output = model(input_tensor) - 可以用
torch.backends.mps.is_available()确认当前环境是否支持MPS设备,避免环境配置问题。
内容的提问来源于stack exchange,提问作者Loman James
相关产品推荐
相关产品推荐

