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
相关产品推荐
相关产品推荐

