PyTorch开发贪吃蛇DQN报错:期望Long标量类型却得到Float求助
问题根源
PyTorch 中 nn.Linear 全连接层的权重、偏置参数默认初始化类型为 torch.float32,而矩阵乘法运算要求输入张量与层参数张量的标量类型完全一致才能运算。
你在模型 forward 方法中主动将输入强制转换为 Long 长整型,同时原始输入从整数数组转张量时默认生成的也是整型张量,和全连接层的浮点型参数类型不匹配,触发了类型错误。此处报错提示的描述存在反向误导,实际是层参数为浮点型,输入为整型,二者类型不兼容。
解决方案
按以下步骤修改即可解决问题:
- 删除模型
forward方法中的input = input.long()语句,不要强制将输入转换为长整型 - 将输入
state转换为和层参数一致的浮点型,修改推理阶段的代码如下:
# EXPLOIT else: state = torch.tensor(state).float() state = state.unsqueeze(0) action_values = self.net(state, model="online") dir = torch.argmax(action_values, axis=1).item()
额外优化建议
你当前定义的网络结构中,所有全连接层之间没有添加非线性激活函数,多层线性层叠加等价于单层线性层,会严重限制模型的拟合能力,建议在每两层 nn.Linear 之间添加 nn.ReLU() 激活函数,修改后的网络结构参考如下:
self.online = nn.Sequential( nn.Linear(input_dim, 200), nn.ReLU(), nn.Linear(200, 20), nn.ReLU(), nn.Linear(20, 50), nn.ReLU(), nn.Linear(50, output_dim), )
内容的提问来源于stack exchange,提问作者gavinjessey
相关产品推荐
相关产品推荐

