You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.27 06:42:41