基于DQN与TF处理嵌套spaces.Dict可变大小观测空间的问题
解决方案方向指引
一、观测格式扁平化转换
Keras/TensorFlow模型无法直接处理嵌套字典+元组的混合结构,必须先把观测转换成一维张量或模型可识别的扁平结构。以下是具体转换逻辑:
- 离散型特征(elem_1、elem_4、elem_3下的所有prop):类别数少则转one-hot编码,类别数多(如elem_1的64类)用**整数嵌入(Embedding)**更高效。
- 连续型特征(elem_2):直接保留为浮点张量。
- elem_3的100个重复Dict元组:将每个Dict的4个prop分别编码后,按顺序拼接成一个长向量。
示例预处理函数:
import numpy as np from tensorflow.keras.utils import to_categorical def preprocess_observation(self, observation): # 处理elem_1:转one-hot elem_1 = to_categorical(observation["elem_1"], num_classes=64) # 处理elem_2:直接转张量 elem_2 = np.array(observation["elem_2"], dtype=np.float32) # 处理elem_3:遍历100个Dict,每个prop转one-hot后拼接 elem_3_list = [] for sub_dict in observation["elem_3"]: prop1 = to_categorical(sub_dict["prop_1"], num_classes=100) prop2 = to_categorical(sub_dict["prop_2"], num_classes=10) prop3 = to_categorical(sub_dict["prop_3"], num_classes=4) prop4 = to_categorical(sub_dict["prop_4"], num_classes=4) elem_3_list.append(np.concatenate([prop1, prop2, prop3, prop4])) elem_3 = np.concatenate(elem_3_list) # 处理elem_4:转one-hot elem_4 = to_categorical(observation["elem_4"], num_classes=32) # 拼接所有特征成一维张量,增加batch维度(model.predict需要批量输入) flattened_obs = np.concatenate([elem_1, elem_2, elem_3, elem_4]) return np.expand_dims(flattened_obs, axis=0)
修改epsilon_greedy_action_selection函数:
def epsilon_greedy_action_selection(self, model, epsilon, observation): logger.info("OBSERVATION RECEIVED : {}".format(len(observation))) type = None if np.random.random() > epsilon: # 先预处理观测 processed_obs = self.preprocess_observation(observation) prediction = model.predict(processed_obs, verbose=0) action = np.argmax(prediction) type = "IF" else: action = random.sample(range(self.action_space.n), 1)[0] type = "ELSE" logger.info(f"Returning action as {action} from {type}") return action
二、构建多输入模型匹配结构化观测
如果不想扁平化特征(担心丢失结构化信息),可以构建多输入模型,每个输入对应观测空间的一个元素:
from tensorflow.keras import Model from tensorflow.keras.layers import Input, Embedding, Flatten, Dense, Concatenate, TimeDistributed def build_dqn_model(self): # 输入1:elem_1(离散64类) elem1_input = Input(shape=(), name="elem_1") elem1_emb = Embedding(input_dim=64, output_dim=16)(elem1_input) elem1_out = Flatten()(elem1_emb) # 输入2:elem_2(连续1维) elem2_input = Input(shape=(1,), name="elem_2") elem2_out = Dense(16, activation="relu")(elem2_input) # 输入3:拆分elem_3的每个prop为序列输入 prop1_input = Input(shape=(100,), name="prop_1") prop2_input = Input(shape=(100,), name="prop_2") prop3_input = Input(shape=(100,), name="prop_3") prop4_input = Input(shape=(100,), name="prop_4") # 对每个prop序列做嵌入+时间分布处理 prop1_emb = Embedding(100, 8)(prop1_input) prop1_out = TimeDistributed(Flatten())(prop1_emb) prop2_emb = Embedding(10, 4)(prop2_input) prop2_out = TimeDistributed(Flatten())(prop2_emb) prop3_emb = Embedding(4, 2)(prop3_input) prop3_out = TimeDistributed(Flatten())(prop3_emb) prop4_emb = Embedding(4, 2)(prop4_input) prop4_out = TimeDistributed(Flatten())(prop4_emb) # 拼接每个prop的输出,再做全连接压缩 elem3_concat = Concatenate(axis=-1)([prop1_out, prop2_out, prop3_out, prop4_out]) elem3_out = Dense(32, activation="relu")(elem3_concat) elem3_out = Flatten()(elem3_out) # 输入4:elem_4(离散32类) elem4_input = Input(shape=(), name="elem_4") elem4_emb = Embedding(32, 8)(elem4_input) elem4_out = Flatten()(elem4_emb) # 拼接所有输入的输出,构建DQN主体网络 concat = Concatenate()([elem1_out, elem2_out, elem3_out, elem4_out]) hidden = Dense(64, activation="relu")(concat) hidden = Dense(32, activation="relu")(hidden) # 输出层:匹配动作空间维度 output = Dense(self.action_space.n, activation="linear")(hidden) # 定义多输入模型 model = Model( inputs=[elem1_input, elem2_input, prop1_input, prop2_input, prop3_input, prop4_input], outputs=output ) model.compile(optimizer="adam", loss="mse") return model
此时预处理观测需要拆分成对应输入的列表:
def preprocess_observation_for_multi_input(self, observation): # 提取elem_3的每个prop序列 prop1_seq = np.array([sub["prop_1"] for sub in observation["elem_3"]]) prop2_seq = np.array([sub["prop_2"] for sub in observation["elem_3"]]) prop3_seq = np.array([sub["prop_3"] for sub in observation["elem_3"]]) prop4_seq = np.array([sub["prop_4"] for sub in observation["elem_3"]]) # 返回模型需要的输入列表(对应模型inputs顺序,增加batch维度) return [ np.array([observation["elem_1"]]), np.array([observation["elem_2"]]), np.array([prop1_seq]), np.array([prop2_seq]), np.array([prop3_seq]), np.array([prop4_seq]) ]
三、简化观测空间(业务允许时)
如果elem_3的100个Dict存在冗余或可聚合信息,可提取统计特征压缩维度:
- 统计prop_1的最大值、最小值、平均值
- 统计prop_2、prop_3、prop_4的类别出现次数
把100个Dict的信息压缩成几十维统计特征,大幅降低输入复杂度,模型更容易处理。
内容的提问来源于stack exchange,提问作者Ai_Nebula
相关产品推荐
相关产品推荐

