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

基于Keras-RL的DQN适配Dict型观测空间输入不匹配报错排查

解决MtEnv Dict观测空间下Keras DQN的输入不匹配问题

问题根源在于你用的Keras DQN实现(比如keras-rl2)默认不支持Dict类型的观测输入,而Stable Baselines3的A2C/PPO内部已经做了Dict观测的解析与适配逻辑。以下是具体解决思路:

方法1:构建支持Dict输入的Keras模型

直接把模型的输入定义为与观测空间对应的字典结构,让模型自动解析每个键的特征:

from tensorflow.keras import Input, Model
from tensorflow.keras.layers import Dense, Concatenate

# 根据MtEnv的observation_space定义各输入层
input_layers = {}
for key in env.observation_space.spaces:
    input_layers[key] = Input(shape=env.observation_space[key].shape, name=key)

# 对每个输入做特征提取(可根据需求调整网络结构)
processed_layers = []
for key in input_layers:
    x = Dense(32, activation='relu')(input_layers[key])
    # 针对不同输入可加专属处理层,比如features维度大就加更多层
    if key == 'features':
        x = Dense(64, activation='relu')(x)
    processed_layers.append(x)

# 合并所有特征并输出动作值
merged = Concatenate()(processed_layers)
x = Dense(128, activation='relu')(merged)
output = Dense(env.action_space.n, activation='linear')(x)

# 构建以字典为输入的模型
model = Model(inputs=input_layers, outputs=output)

训练时直接传入观测字典即可,不需要额外拆解。如果用keras-rl2的DQN Agent,需要确保回放缓冲区存储完整的观测字典,采样后直接喂给模型。

方法2:包装环境,将Dict观测拆分为多输入数组

如果不想修改模型结构,可以创建环境包装器,把Dict观测转换成模型期望的多个输入数组:

import gym
from gym import Wrapper

class DictToMultiInputWrapper(Wrapper):
    def __init__(self, env):
        super().__init__(env)
        # 这里可以根据模型输入顺序定义键的列表
        self.input_keys = ['balance', 'equity', 'margin', 'features', 'orders']

    def reset(self, **kwargs):
        obs_dict = self.env.reset(**kwargs)
        # 按顺序提取各输入数组
        return tuple(obs_dict[key] for key in self.input_keys)

    def step(self, action):
        obs_dict, reward, done, info = self.env.step(action)
        obs_tuple = tuple(obs_dict[key] for key in self.input_keys)
        return obs_tuple, reward, done, info

# 包装原环境
wrapped_env = DictToMultiInputWrapper(env)

之后用包装后的环境训练DQN,此时模型接收的就是拆分后的多个输入数组,与模型的输入层顺序对应即可。

方法3:修改训练循环的数据处理逻辑

如果用自定义训练循环,手动在数据采样阶段拆解观测字典:

# 假设replay_buffer存储的是(obs_dict, action, reward, next_obs_dict, done)样本
batch = replay_buffer.sample(batch_size)
obs_batch = {key: np.array([sample[0][key] for sample in batch]) for key in input_keys}
next_obs_batch = {key: np.array([sample[3][key] for sample in batch]) for key in input_keys}
action_batch = np.array([sample[1] for sample in batch])
reward_batch = np.array([sample[2] for sample in batch])
done_batch = np.array([sample[4] for sample in batch])

# 喂给模型训练
model.train_on_batch(obs_batch, target_q_values)

这种方式灵活度高,适合需要自定义训练流程的场景。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.21 12:33:16