基于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
相关产品推荐
相关产品推荐

