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赋值,再判断游戏是否结束,状态赋值完全不符合预期
修复点如下:
- 移除
done = True的硬编码,将done的赋值放到游戏结束判断处 - 调大经验回放池容量,建议至少设置为10000,远大于BATCH_SIZE
- 调整状态赋值与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
相关产品推荐
相关产品推荐

