2048游戏DQN实现中Byte与Float dtype不匹配问题求助
DQN实现2048游戏时的矩阵乘法数据类型不匹配问题解决与调试指南
问题核心
在基于PyTorch实现2048游戏的DQN智能体时,触发了矩阵乘法数据类型不匹配错误:一方为Byte类型张量,另一方为Float类型。尝试在报错代码后追加.float()转换无效,同时遇到PyCharm调试定位困难的问题。
报错关联代码
next_state_values[non_final_mask] = target_net(non_final_next_states).max(1).values x = F.relu(self.layer1(x))
相关代码片段
DQN网络类
class DQN(nn.Module): def __init__(self, n_observations, n_actions): super(DQN, self).__init__() self.layer1 = nn.Linear(16, 256) self.layer2 = nn.Linear(256, 256) self.layer3 = nn.Linear(256, n_actions) def forward(self, x): x = F.relu(self.layer1(x)).float() x = F.relu(self.layer2(x)) return self.layer3(x)
模型优化函数
def optimize_model(): if len(memory) < BATCH_SIZE: return transitions = memory.sample(BATCH_SIZE) batch = Transition(*zip(*transitions)) non_final_mask = torch.tensor(tuple(map(lambda s: s is not None, batch.n_state)), dtype=torch.bool) non_final_next_states = torch.tensor([s for s in batch.n_state if s is not None]) state_batch = torch.stack([torch.tensor(s) for s in batch.state]) action_batch = torch.stack([torch.tensor(s) for s in batch.action]).unsqueeze(1) reward_batch = torch.stack([torch.tensor(s) for s in batch.reward]) state_action_values = policy_net(state_batch).gather(1, action_batch) next_state_values = torch.zeros(BATCH_SIZE) with torch.no_grad(): next_state_values[non_final_mask] = target_net(non_final_next_states).max(1).values.float() expected_state_action_values = (next_state_values * DISCOUNT_FACTOR) + reward_batch criterion = nn.SmoothL1Loss() loss = criterion(state_action_values, expected_state_action_values.unsqueeze(1)) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_value_(policy_net.parameters(), 100) optimizer.step()
解决与调试方案
1. 根治类型不匹配问题
问题根源是输入到网络的状态张量为Byte类型,而网络层参数默认是Float类型,需从源头转换:
- 修改状态张量构造代码,强制转为Float32:
# 修正state_batch state_batch = torch.stack([torch.tensor(s, dtype=torch.float32) for s in batch.state]) # 修正non_final_next_states non_final_next_states = torch.tensor([s for s in batch.n_state if s is not None], dtype=torch.float32) # 修正reward_batch(避免后续运算类型冲突) reward_batch = torch.stack([torch.tensor(s, dtype=torch.float32) for s in batch.reward]) - 移除
forward中多余的.float():线性层输出本身就是Float类型,额外转换会引发类型混乱,将x = F.relu(self.layer1(x)).float()改为x = F.relu(self.layer1(x))。
2. 高效调试定位问题
- 关键位置打印类型:在
optimize_model处理完批次数据后、forward函数开头添加打印语句,快速定位异常张量:# 在optimize_model中 print("state_batch dtype:", state_batch.dtype) print("non_final_next_states dtype:", non_final_next_states.dtype) # 在DQN.forward中 def forward(self, x): print("Input x dtype:", x.dtype) # ... 其余代码 - 条件断点过滤无效触发:在PyCharm中给
forward函数第一行设置条件断点,条件设为x.dtype != torch.float32,仅当输入类型异常时触发断点,避开大量正常过渡状态的干扰。
3. 验证修复效果
运行代码前,确保所有输入到网络的张量(state_batch、non_final_next_states)均为Float32类型,网络层参数类型与输入一致,即可解决矩阵乘法类型不匹配错误。
内容的提问来源于stack exchange,提问作者Elly Sinden
相关产品推荐
相关产品推荐

