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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 17:40:23