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

Keras RL2 DQAgent训练FlappyBird环境时维度不匹配报错求助

Keras RL2 DQNAgent训练自定义FlappyBird环境时维度不匹配报错排查

ValueError: Error when checking input: expected dense_input to have 2 dimensions, but got array with shape (1, 3, 4)

错误发生在dqnagent.fit执行过程中,调用栈指向rl/core.py第168行的action = self.forward(observation)。已确认自定义环境的Game()模块无问题(曾用于NEAT项目)。

环境代码

class flappy_env(Env):

    def __init__(self):
        self.game = Game()
        self.observation_space = Box(low=np.array([-0.4, -2.0, -1.0, -1.0], dtype=np.float32),
                                     high=np.array([1.0, 2.0, 1.0, 0.5], dtype=np.float32))
        self.action_space = Discrete(2)

    def step(self, action):
        done, score, reward = self.game.play_step(action)
        state = self.game.get_state()
        info = {}
        return state, reward, done, info

    def render(self):
        pass

    def reset(self):
        self.game.reset_game()
        return self.game.get_state()

状态获取代码

def get_state(self):
    if len(self.pipe_group) > 0:
       bird_y_loc = self.flappy.rect.y
       x_dist_pipe_bird = self.pipe_group.sprites()[0].rect.left - self.flappy.rect.right
       bot_pipe_y_loc = self.pipe_group.sprites()[0].rect.top - bird_y_loc
       top_pipe_y_loc = self.pipe_group.sprites()[1].rect.bottom - bird_y_loc
       return np.array([x_dist_pipe_bird / 500, 10 * bot_pipe_y_loc / screen_height,
                        5 * top_pipe_y_loc / screen_height, self.flappy.vel / 35], dtype=np.float32)
    # shouldn't get here
    return None

模型构建代码

from keras.models import Sequential
from keras.layers import Dense
from keras.optimizers import Adam
from rl.agents import DQNAgent
from rl.memory import SequentialMemory
from rl.policy import LinearAnnealedPolicy, EpsGreedyQPolicy
from env import flappy_env


def build_model():
    model = Sequential()
    model.add(Dense(16, input_shape=(4,), activation='relu'))
    model.add(Dense(16, activation='relu'))
    model.add(Dense(2, activation='linear'))
    return model


def build_agent(sequential_model):
    policy = LinearAnnealedPolicy(EpsGreedyQPolicy(), attr='eps', value_max=1., value_min=.1, value_test=.2, nb_steps=10000)
    memory = SequentialMemory(limit=1000, window_length=3)
    dqn = DQNAgent(model=sequential_model, memory=memory, policy=policy,
                   enable_dueling_network=True, dueling_type='avg',
                   nb_actions=2, nb_steps_warmup=1000)
    return dqn


env = flappy_env()
model = build_model()
model.summary()
dqn = build_agent(model)
dqn.compile(Adam(learning_rate=1e-4))
dqn.fit(env, nb_steps=10000, visualize=False)

调试后确认报错中的维度3对应SequentialMemory的window_length参数,dqn.py中q_values = self.compute_batch_q_values([state]).flatten()代码会添加额外维度,期望状态不被添加额外维度,以下是排查修复方案:


修复方案

问题核心是SequentialMemory的window_length参数与模型输入形状不匹配:当设置window_length=3时,Keras RL2会自动拼接连续3个时间步的状态,形成(window_length, state_dim)的输入形状,但你的模型仅定义了单个状态维度(4,),导致维度冲突。

有两种解决思路:

  1. 适配多步状态输入
    修改模型结构,让其能处理3个连续状态的输入:

    from keras.layers import Flatten
    
    def build_model():
        model = Sequential()
        model.add(Flatten(input_shape=(3,4)))  # 展平3*4的多步状态
        model.add(Dense(16, activation='relu'))
        model.add(Dense(16, activation='relu'))
        model.add(Dense(2, activation='linear'))
        return model
    

    此时模型可接受(1,3,4)形状的输入(批量维度+时间步+状态维度)。

  2. 取消多步状态拼接
    若不需要使用连续多步状态,直接将window_length改为1:

    memory = SequentialMemory(limit=1000, window_length=1)
    

    模型输入(4,)与单个状态形状完全匹配,不会再出现维度问题。

另外注意:原build_agent函数中memory = ...一行存在缩进错误,需与policy = ...保持同一层级,避免语法报错。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 03:23:22