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

PyTorch实现Deep Q-Learning时空张量及卷积通道不匹配报错求助

问题排查与修复方案

第二类通道不匹配报错修复

报错的核心是DQN网络结构定义错误,存在两处参数不匹配:

  • 第一层卷积conv1输出通道为32,第二层卷积conv2的输入通道错误设置为16,需要修改为32
  • 第三层卷积conv3输出通道为32,计算全连接层输入维度时错误乘以1,需要修改为乘以32
    修复后的DQN类核心代码如下:
class DQN(nn.Module):
    def __init__(self, h, w, outputs):
        super(DQN, self).__init__()
        self.conv1 = nn.Conv2d(3, 32, kernel_size=5,stride=2 )
        self.bn1 = nn.BatchNorm2d(32)
        # 修改变卷积2输入通道为32
        self.conv2 = nn.Conv2d(32, 32, kernel_size=5, stride=2)
        self.bn2 = nn.BatchNorm2d(32)
        self.conv3 = nn.Conv2d(32, 32, kernel_size=5, stride=2)
        self.bn3 = nn.BatchNorm2d(32)

        def conv2d_size_out(size, kernel_size = 5, stride = 2):
            return (size - (kernel_size - 1) - 1) // stride  + 1
        convw = conv2d_size_out(conv2d_size_out(conv2d_size_out(w)))
        convh = conv2d_size_out(conv2d_size_out(conv2d_size_out(h)))
        # 修改全连接输入维度计算,乘以32
        linear_input_size = convw * convh * 32
        self.head = nn.Linear(linear_input_size, outputs)

    def forward(self, x):
        x = F.relu(self.bn1(self.conv1(x)))
        x = F.relu(self.bn2(self.conv2(x)))
        x = F.relu(self.bn3(self.conv3(x)))
        return self.head(x.view(x.size(0), -1))

第一类torch.cat空列表报错修复

报错由逻辑错误导致,存在三处问题:

  • 主逻辑中硬编码done = True,导致每一步的next_state都被设为None,采样的batch中没有有效next_state,拼接空列表触发报错
  • 经验回放池ReplayMemory容量仅设置为50,当BATCH_SIZE≥128时永远无法达到执行优化的阈值,优化逻辑永不触发;调低BATCH_SIZE后触发优化时,所有next_state都是None,就会触发空拼接错误
  • done的判断逻辑位置错误,当前逻辑是先给done赋值,再判断游戏是否结束,状态赋值完全不符合预期
    修复点如下:
  1. 移除done = True的硬编码,将done的赋值放到游戏结束判断处
  2. 调大经验回放池容量,建议至少设置为10000,远大于BATCH_SIZE
  3. 调整状态赋值与done的判断顺序
    修复后的主逻辑核心代码如下:
# 调大回放池容量
memory = ReplayMemory(10000)

num_episodes = 50
for i_episode in range(num_episodes):
    last_screen = gym.take_screen("screen/shot.jpg")
    current_screen = gym.take_screen("screen/shot.jpg")
    state = current_screen - last_screen
    for t in count():
        action = select_action(state)
        getattr(gym, gym.actions[action.item()])()
        reward = torch.tensor([gym.last_update_on_score], device=device)
        # 先判断游戏是否结束
        done = gym.get_score() and gym.last_score()
        # 再处理状态
        last_screen = current_screen
        current_screen = gym.take_screen("screen/shot.jpg")
        if not done:
            next_state = current_screen - last_screen
        else:
            next_state = None
            print("gameover")
            gym.game_quit()
        
        memory.push(state, action, next_state, reward)
        state = next_state
        optimize_model()
        if done:
            episode_durations.append(t + 1)
            break
    if i_episode % TARGET_UPDATE == 0:
        target_net.load_state_dict(policy_net.state_dict())

内容的提问来源于stack exchange,提问作者SombreroMickey

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 15:48:04