PyTorch MPS后端调用from_numpy()报is_mps()校验失败问题
报错根因
触发该错误的核心原因是模型与输入张量所在计算设备不匹配:
- 执行
model.to(device)后,模型的全部权重参数已经迁移到MPS设备,所有前向计算逻辑要求输入张量必须同样驻留在MPS设备内存中 torch.from_numpy()方法生成的张量默认存储在CPU内存,不会自动跨设备同步到MPS,直接传入MPS上的模型时,线性层计算前的设备校验就会失败,抛出对应RuntimeError- 切换到CPU运行无报错,是因为此时模型权重和输入张量都在CPU内存,设备完全匹配。
修复方案
在numpy数组转PyTorch张量后,显式将张量迁移到和模型一致的运行设备即可,修正后的推理代码如下:
observations = env.reset() # numpy转张量后,同步迁移到MPS/CPU对应设备,可同时指定数据类型对齐模型要求 X = torch.from_numpy(observations).to(device=device, dtype=torch.float32) logits = model(X)
注意:PyTorch不会自动做跨设备的张量拷贝,不管是MPS还是CUDA后端,只要模型运行在非CPU设备上,所有输入张量都需要显式调用
.to(device)迁移到对应设备后再传入模型。
内容的提问来源于stack exchange,提问作者sandboxj
相关产品推荐
相关产品推荐

