PyTorch DQN训练Gymnasium俄罗斯方块遇类型不匹配RuntimeError
问题排查与解决
错误原因
报错Input type (unsigned char) and bias type (float) should be the same的核心原因是:输入模型的state张量类型为unsigned char(uint8),而DQN模型的参数(如偏置项)是float类型,PyTorch不允许不同数据类型的张量进行运算。
Gymnasium的俄罗斯方块环境返回的观测值(即state)通常是0-255的像素数据,默认以uint8类型存储,这和模型默认的float32参数类型不兼容。
解决方案
只需要在将state传入模型前,将其转换为与模型参数一致的浮点类型即可,同时确保设备(CPU/CUDA)与模型一致:
代码修改示例
在执行action_values = self.net(state, model="current")前添加类型转换代码:
# 转换state为float32类型 state = state.float() # 若使用CUDA,确保state与模型在同一设备上(需提前定义模型的device属性) state = state.to(self.net.device) # 再传入模型计算 action_values = self.net(state, model="current")
或者在获取环境state时直接处理:
# 从环境获取state后立即转换类型并移到目标设备 state, info = env.reset() state = torch.tensor(state, dtype=torch.float32).unsqueeze(0).to(self.net.device)
验证方法
可以在转换前打印state.dtype确认原类型,转换后再次打印确认类型变为torch.float32,确保和模型参数(如self.net.fc1.bias.dtype)一致。
内容的提问来源于stack exchange,提问作者AJ Andersen
相关产品推荐
相关产品推荐

