Gym数独环境用KerasRL训练时输入维度不匹配报错
解决KerasRL训练数独环境时的输入维度不匹配问题
问题原因
报错显示模型期望3维输入,但实际接收到(1,1,9,9)的4维数组。这是因为KerasRL的DQNAgent结合SequentialMemory时,会自动在状态前添加window_length维度(你设置的是1),再加上batch维度,最终输入变成(batch_size, window_length, 9,9),而你的模型输入层定义的是(9,9),维度不匹配导致报错。
修复步骤
1. 修改模型输入层适配window_length
在build_model函数中,调整输入形状以包含window_length维度,让模型能接收KerasRL传递的4维输入:
def build_model(states, actions, window_length=1): model = Sequential() # 输入形状改为 (window_length, 9, 9) model.add(Input(shape=(window_length,) + states)) model.add(Flatten()) # 将1*9*9的输入展平为81维向量 model.add(Dense(100, activation='relu')) model.add(Dense(100, activation='relu')) model.add(Dense(actions, activation='linear')) return model
2. 调用模型时传入window_length参数
创建模型时,同步传入DQNAgent中设置的window_length值:
window_length = 1 model = build_model(states, actions, window_length)
3. 额外提示:完善数独环境核心逻辑
你的Sudoku环境目前仅为框架,step函数未实现动作解析、状态更新、奖励计算和终止条件,render函数也未完成,这些会导致训练无实际意义,后续需要补充:
- 在
step函数中解析729种动作(对应81个格子×9个数字),更新棋盘状态 - 根据数独规则设计奖励机制(如填对数字加正奖励、填错加负奖励、完成数独加大额奖励)
- 设置
done标志(数独完成或无法继续时返回True)
修改后的核心代码片段
def build_model(states, actions, window_length=1): model=Sequential() model.add(Input(shape=(window_length,) + states)) model.add(Flatten()) model.add(Dense(100,activation='relu')) model.add(Dense(100,activation='relu')) model.add(Dense(actions,activation='linear')) return model # 其余代码不变 env=Sudoku() states=env.observation_space.shape print(states) actions=env.action_space.n window_length = 1 model = build_model(states,actions,window_length) model.summary() agent = build_agent(model,actions) agent.compile(Adam(lr=1e-3),metrics=['mae']) agent.fit(env,nb_steps=20000,visualize=False,verbose=1)
内容的提问来源于stack exchange,提问作者Lucas 'Snufkin' Gautier
相关产品推荐
相关产品推荐

