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

如何让Gym.Env输出两个独立数组适配DQN神经网络

问题:Gym环境输出无法匹配DQN多输入网络

我需要让Gym.Env输出两个独立数组,分别传入DQN神经网络,但它总是把两个数组合并成一个包含两个数组的单个数组,导致模型输入不匹配。已经尝试调整observation_space使用Box和Tuple,仍未解决问题。

相关代码

class GoEnv(gym.Env):

    def __init__(self):
        self.action_space = spaces.Discrete(3)
        self.observation_space = spaces.Tuple([spaces.Box(low=-np.inf, high=np.inf, shape=(2, 11), dtype=np.float32),
                                               spaces.Box(low=-np.inf, high=np.inf, shape=(1, 11), dtype=np.float32)])

    def step(self, action):
        state = [np.array(self.data), np.array(self.account)]
        return state, reward, self.done, info

envi = env.GoEnv()

def data_model():
    data_input = layers.Input(shape=(500, 2, 11))
    acc_input = layers.Input(shape=(500, 1, 11))

    dat_model = layers.Conv2D(filters=32, activation='swish', kernel_size=(500, 1),
                              padding='valid', strides=(500, 1))(data_input)
    dat_model = layers.Dense(3, activation='swish')(dat_model)
    dat_model = layers.Dense(3, activation='softmax')(dat_model)
    dat_model = layers.Flatten()(dat_model)
    dat_model = keras.Model(inputs=data_input, outputs=dat_model)

    acc_model = layers.Dense(3, activation='swish')(acc_input)
    acc_model = layers.Dense(3, activation='softmax')(acc_model)
    acc_model = layers.Flatten()(acc_model)
    acc_model = keras.Model(inputs=acc_input, outputs=acc_model)

    combined = layers.concatenate([dat_model.output, acc_model.output])

    z = layers.Flatten()(combined)
    z = layers.Dense(64, activation='swish')(z)
    z = layers.Dense(3, activation='softmax')(z)

    model = keras.Model(inputs=[dat_model.input, acc_model.input], outputs=z)

    return model

model = data_model()
model.summary()
actions = 3

def build_agent(model, actions):
    policy = BoltzmannQPolicy()
    memory = SequentialMemory(limit=50000, window_length=500)
    dqn = DQNAgent(model=model,
                   memory=memory,
                   policy=policy,
                   nb_actions=actions,
                   nb_steps_warmup=600,
                   target_model_update=1e-2)
    return dqn
dqn = build_agent(model, actions)
dqn.fit(envi, nb_steps=6000, visualize=False, verbose=1)

报错信息

Traceback (most recent call last):
  File "C:/Users/Worrall/PycharmProjects/Prject/main.py", line 46, in <module>
    dqn.fit(envi, nb_steps=6000, visualize=False, verbose=1)
  File "C:\Users\Worrall\PycharmProjects\DocumentRecog\venv\lib\site-packages\rl\core.py", line 168, in fit
    action = self.forward(observation)
  File "C:\Users\Worrall\PycharmProjects\DocumentRecog\venv\lib\site-packages\rl\agents\dqn.py", line 224, in forward
    q_values = self.compute_q_values(state)
  File "C:\Users\Worrall\PycharmProjects\DocumentRecog\venv\lib\site-packages\rl\agents\dqn.py", line 68, in compute_q_values
    q_values = self.compute_batch_q_values([state]).flatten()
  File "C:\Users\Worrall\PycharmProjects\DocumentRecog\venv\lib\site-packages\rl\agents\dqn.py", line 63, in compute_batch_q_values
    q_values = self.model.predict_on_batch(batch)
  File "C:\Users\Worrall\PycharmProjects\DocumentRecog\venv\lib\site-packages\tensorflow\python\keras\engine\training_v1.py", line 1200, in predict_on_batch
    inputs, _, _ = self._standardize_user_data(
  File "C:\Users\Worrall\PycharmProjects\DocumentRecog\venv\lib\site-packages\tensorflow\python\keras\engine\training_v1.py", line 2328, in _standardize_user_data
    return self._standardize_tensors(
  File "C:\Users\Worrall\PycharmProjects\DocumentRecog\venv\lib\site-packages\tensorflow\python\keras\engine\training_v1.py", line 2356, in _standardize_tensors
    x = training_utils.standardize_input_data(
  File "C:\Users\Worrall\PycharmProjects\DocumentRecog\venv\lib\site-packages\tensorflow\python\keras\engine\training_utils.py", line 533, in standardize_input_data
    raise ValueError('Error when checking model ' + exception_prefix +
ValueError: 检查模型输入时出错:传入模型的Numpy数组列表大小不符合预期。预期看到2个数组,对应输入['input_1', 'input_2'],但实际收到1个数组:[array([[[array([[0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.],
       [0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.]]),
         array([[0., 0., 0., 0., 0., 0., 0., 0., 0., 0., 0.]])],
        [array([[...

问题分析

报错核心是模型期望2个独立输入数组,但实际只收到1个嵌套数组。原因包括:

  • Gym的Tuple观测空间返回的结果,会被Keras-RL自动打包成一个嵌套数组,而非两个独立输入
  • SequentialMemory设置的window_length=500会堆叠观测,多输入场景下需要手动拆分堆叠后的批量数据

解决方案

1. 修正环境的观测输出格式

将step方法返回的观测改为元组类型的独立numpy数组,确保维度和数据类型正确:

def step(self, action):
    # 确保self.data是(2,11)、self.account是(1,11)的float32类型数组
    state = (np.array(self.data, dtype=np.float32), np.array(self.account, dtype=np.float32))
    return state, reward, self.done, info

2. 自定义多输入处理器

实现Processor子类,拆分堆叠后的批量观测,适配模型输入要求:

from rl.core import Processor

class MultiInputProcessor(Processor):
    def process_state_batch(self, batch):
        # 拆分批量数据为两个独立的输入组
        batch_data = np.array([sample[:, 0] for sample in batch])
        batch_acc = np.array([sample[:, 1] for sample in batch])
        # 调整维度匹配模型输入:(样本数, window_length, 2, 11) 和 (样本数, window_length, 1, 11)
        batch_data = batch_data.reshape(batch_data.shape[0], 500, 2, 11)
        batch_acc = batch_acc.reshape(batch_acc.shape[0], 500, 1, 11)
        return [batch_data, batch_acc]

3. 修改Agent初始化逻辑

在构建DQN Agent时添加自定义处理器:

def build_agent(model, actions):
    policy = BoltzmannQPolicy()
    memory = SequentialMemory(limit=50000, window_length=500)
    processor = MultiInputProcessor()
    dqn = DQNAgent(model=model,
                   memory=memory,
                   policy=policy,
                   nb_actions=actions,
                   nb_steps_warmup=600,
                   target_model_update=1e-2,
                   processor=processor)
    return dqn

4. 验证输入维度一致性

确保process_state_batch返回的数组形状,与模型输入层定义的shape完全匹配:

  • data_input的shape=(500,2,11)对应处理后的batch_data
  • acc_input的shape=(500,1,11)对应处理后的batch_acc

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.17 01:46:47