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

如何解决PyTorch中Atari强化学习项目的尺寸不匹配错误?

解决PyTorch Atari强化学习模型尺寸不匹配问题

问题根源

错误提示的尺寸不匹配,是因为当前模型中全连接层7.weight的输入维度(409600)与checkpoint中预训练模型的对应维度(3136)不一致。3136对应标准DQN架构中卷积层输出的64*7*7特征图,你的代码问题出在:

  • 卷积层的padding计算错误,导致卷积后特征图尺寸过大;
  • 输入图像可能未做标准预处理(如缩放到84x84),或全连接层输入维度未对应标准特征图尺寸。

修改方案

1. 对齐标准DQN卷积层参数

标准Atari DQN的卷积层参数是固定的,直接使用以下设置即可得到64*7*7的特征图输出,与checkpoint维度匹配:

self.q_network = nn.Sequential(
    # 输入为堆叠的state_history帧(通常4帧),预处理后84x84灰度图
    nn.Conv2d(in_channels=n_channels * self.config["hyper_params"]["state_history"], 
              out_channels=32, kernel_size=8, stride=4),
    nn.ReLU(),
    nn.Conv2d(in_channels=32, out_channels=64, kernel_size=4, stride=2),
    nn.ReLU(),
    nn.Conv2d(in_channels=64, out_channels=64, kernel_size=3, stride=1),
    nn.ReLU(),
    nn.Flatten(),
    # 对应64*7*7=3136,匹配checkpoint的输入维度
    nn.Linear(64 * 7 * 7, 512),
    nn.ReLU(),
    nn.Linear(in_features=512, out_features=num_actions)
)

2. 确保输入图像预处理正确

必须将原始Atari游戏画面做如下预处理,否则即使修改卷积层参数仍会出错:

  • 转换为灰度图(单通道);
  • 缩放到84x84尺寸;
  • 堆叠state_history帧(通常为4帧)作为模型输入。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 02:02:47