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

基于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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 02:02:06