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

FlappyBird卷积DQN训练报Input_input维度不匹配错误

问题报错核心原因

触发报错的核心日志:ValueError: Error when checking input: expected Input_input to have 4 dimensions, but got array with shape (1, 1, 2)

  • 输入数据与网络层要求的维度、结构完全不匹配:Conv2D层要求输入为4维张量(按channels_first格式为[批次大小, 通道数, 特征高度, 特征宽度],按channels_last格式为[批次大小, 特征高度, 特征宽度, 通道数]),但你使用的flappy_bird_gym的FlappyBird-v0环境默认返回的观测是形状为(2,)的一维数值向量,两个值分别对应小鸟到下一根管道的水平距离、小鸟与管道缺口的垂直高度差,根本不存在卷积层需要的二维空间结构。
  • 网络输入定义逻辑错误:你给首层Conv2D设置的input_shape=(1,obs,1),本质是强行给长度为2的一维特征虚构了空间维度,keras-rl的DQNAgent在传入状态时会自动拼接时间步、批次维度,最终传入模型的实际张量形状为(1,1,2),比Conv2D要求的4维输入少1个维度,直接触发维度校验报错。
  • 参考实现适配性偏差:你参考的纯Dense层开源实现可以正常运行,是因为Dense层仅要求最后一维特征数匹配,不需要输入具备空间结构,直接把卷积层套在低维数值观测上本身就不符合卷积层的适用场景。
可落地修改方案

根据你的实际需求二选一即可:

方案1:沿用默认低维数值观测(训练速度快,实现简单)

直接删除所有Conv2D、MaxPooling2D类卷积相关层,改用全连接层搭建网络即可,不需要强行加卷积结构。可参考如下网络定义:

from tensorflow.keras.models import Sequential
from tensorflow.keras.layers import Dense, Flatten

# 对齐环境空间定义
n_actions = env.action_space.n
obs_shape = env.observation_space.shape

model = Sequential()
# keras-rl会自动给输入增加1维时间步维度,这里输入形状对齐为(1, 特征数)
model.add(Flatten(input_shape=(1,) + obs_shape))
model.add(Dense(64, activation='relu'))
model.add(Dense(32, activation='relu'))
# 输出所有动作的Q值
model.add(Dense(n_actions, activation='linear'))

后续按原有流程用keras-rl构建DQNAgent、编译、调用fit()即可正常训练。

方案2:保留Conv2D层,使用屏幕像素作为观测(视觉类RL通用方案)

如果确实需要用卷积层提取视觉特征,就不能用环境默认返回的二维数值状态,需要将输入替换为游戏画面像素,操作步骤如下:

  • 替换环境返回值:不要用默认输出数值状态的接口,自行封装wrapper采集游戏渲染的RGB画面帧,将原始帧处理为84*84分辨率的灰度图,连续堆叠4帧作为单步状态,处理后的单样本状态形状为(84, 84, 4)(对应tensorflow默认的channels_last格式)
  • 调整网络输入定义:首层Conv2D的input_shape设置为和处理后的像素状态一致,即input_shape=(84,84,4),后续再堆叠卷积、池化、Flatten、全连接层即可
  • 增加预处理逻辑:在智能体和环境之间加预处理环节,保证每一步传入模型的状态都经过转灰度、缩放、像素值归一化、帧堆叠处理,维度和网络输入层定义完全匹配
  • 注意:该方案训练算力要求远高于低维数值观测方案,需要GPU加速才能在合理时间内收敛。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.01 23:36:17