PyTorch报错expected scalar type Float but found Double如何解决
问题原因
- PyTorch 中
nn.Linear层默认初始化的权重、偏置参数均为torch.float32类型,也就是报错信息中提到的Float类型 - 你指定
dtype='float64'转换得到的numpy数组转成Torch张量后,对应类型为torch.float64,也就是报错信息中提到的Double类型 - 模型前向计算要求输入张量和模型参数的标量类型完全匹配,二者不一致就会触发该RuntimeError。你将数据转成了双精度浮点,和默认的单精度浮点参数不匹配,所以报错。
修复方案
你可以任选以下任意一种方案解决类型不匹配问题:
方案1:将输入张量转换为单精度浮点(推荐,运算速度更快、显存占用更低)
修改你的张量转换代码即可:
# 方法1:直接指定numpy数组类型为float32 y_2 = torch.from_numpy(np.array(y, dtype='float32')) X_2 = torch.from_numpy(np.array(X, dtype='float32')) # 方法2:转换为Torch张量后调用.float()方法统一转单精度,写法更简洁 y_2 = torch.from_numpy(np.array(y)).float() X_2 = torch.from_numpy(np.array(X)).float()
额外注意:你使用的损失函数
CrossEntropyLoss要求分类标签输入为torch.long(整数类型),如果你的y是分类标签,不需要转浮点,直接转long类型即可,否则后续计算损失还会触发类型错误。
方案2:将模型参数转换为双精度浮点,适配你的输入类型
在定义模型后添加.double()调用,将整个模型的所有参数转为双精度:
net = nn.Linear(54, 7).double()
内容的提问来源于stack exchange,提问作者Ceci Chaung
相关产品推荐
相关产品推荐

