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

FinRL StockTradingEnv状态表示问题:新增特征后维度不匹配

问题:FinRL环境状态空间维度不匹配,新增特征未正确纳入状态初始化

我正在FinRL元环境中调整集成代理的自定义交易策略,添加了神经网络生成的新特征(MYNN)来捕捉数据集的时间动态。处理后数据集的列信息及状态空间计算如下:

INDICATORS = ['macd','rsi_30','cci_30','dx_30','wr_30','atr_30','chop_30','mfi_30','boll_ub','boll_lb','close_30_sma','close_60_sma']

MYNN = ['NN0','NN1','NN2','NN3','NN4']

df_columns = ['date','open','high','low','close','volume','tic','day','macd','rsi_30','cci_30','dx_30','wr_30','atr_30','chop_30','mfi_30','boll_ub','boll_lb','close_30_sma','close_60_sma','vix','turbulence','NN0','NN1','NN2','NN3','NN4']

stock_dimension = len(df_final.tic.unique())
state_space = 1 + 2*stock_dimension + (len(INDICATORS)+len(MYNN ))*stock_dimension
print(f"Stock Dimension: {stock_dimension}, State Space: {state_space}")
#Stock Dimension: 27, State Space: 514

按公式计算状态空间应为514,但环境的_initiate_state函数生成的状态维度仅为442,导致报错:无法将形状(442,)的输入数组广播到形状(514,)。我已经调整了该函数,但仍无法定位问题,目标是让代理能获取新增特征,确保这些特征参与状态初始化。调整后的函数代码如下:

#new_features_list -----> myNN features
def _initiate_state(self):
    if self.initial:
        # For Initial State
        if len(self.df.tic.unique()) > 1:
            # for multiple stock
            state = (
                [self.initial_amount]
                + self.data.close.values.tolist()
                + self.num_stock_shares
                + sum(
                    (self.data[tech].values.tolist() for tech in self.tech_indicator_list),
                    [],
                )
                + sum((self.data[feature].values.tolist() for feature in self.new_features_list), [])
            )
        else:
            # for single stock
            state = (
                [self.initial_amount]
                + [self.data.close]
                + [0] * self.stock_dim
                + sum(([self.data[tech]] for tech in self.tech_indicator_list), [])
                + sum(([self.data[feature]] for feature in self.new_features_list), [])
            )
    else:
        # Using Previous State
        if len(self.df.tic.unique()) > 1:
            # for multiple stock
            state = (
                [self.previous_state[0]]
                + self.data.close.values.tolist()
                + self.previous_state[(self.stock_dim + 1) : (self.stock_dim * 2 + 1)]
                + sum(
                    (self.data[tech].values.tolist() for tech in self.tech_indicator_list),
                    [],
                )
                + sum((self.data[feature].values.tolist() for feature in self.new_features_list), [])
            )
        else:
            # for single stock
            state = (
                [self.previous_state[0]]
                + [self.data.close]
                + self.previous_state[(self.stock_dim + 1) : (self.stock_dim * 2 + 1)]
                + sum(([self.data[tech]] for tech in self.tech_indicator_list), [])
                + sum(([self.data[feature]] for feature in self.new_features_list), [])
            )
    print(f"Initialized state shape: {len(state)} (Expected: {self.state_space})")
    return state

问题定位与修复方案

核心排查方向

状态维度差为514-442=72,说明状态构造过程中某部分元素缺失,按以下步骤逐一验证:

  1. 确认特征列表的完整性
    在_initiate_state函数开头添加打印,检查实际使用的特征数量是否与定义一致:

    print(f"Tech indicators count: {len(self.tech_indicator_list)}")
    print(f"NN features count: {len(self.new_features_list)}")
    

    确保前者长度为12(对应INDICATORS),后者为5(对应MYNN),如果数量不足,需修正tech_indicator_list和new_features_list的赋值逻辑。

  2. 验证状态空间参数已更新
    检查环境初始化代码,确保self.state_space被设置为计算出的514,而非沿用FinRL默认值:

    self.state_space = 1 + 2*stock_dimension + (len(INDICATORS)+len(MYNN ))*stock_dimension
    
  3. 检查持股数列表维度
    多股票场景下,self.num_stock_shares应为长度27的列表(初始全为0),如果是单元素列表会直接少26个维度,添加打印验证:

    print(f"Num stock shares length: {len(self.num_stock_shares)}")
    
  4. 拆分状态构造过程定位缺失
    将状态拆分为各组成部分并打印长度,精准定位哪一块元素不足:

    # 在多股票初始状态分支添加
    part1 = [self.initial_amount]
    part2 = self.data.close.values.tolist()
    part3 = self.num_stock_shares
    part4 = sum((self.data[tech].values.tolist() for tech in self.tech_indicator_list), [])
    part5 = sum((self.data[feature].values.tolist() for feature in self.new_features_list), [])
    print(f"各部分长度:part1={len(part1)}, part2={len(part2)}, part3={len(part3)}, part4={len(part4)}, part5={len(part5)}")
    state = part1 + part2 + part3 + part4 + part5
    

    正常输出应为:part1=1, part2=27, part3=27, part4=324, part5=135,总和514,哪块长度不符就排查对应数据源或列表逻辑。

内容的提问来源于stack exchange,提问作者alee stvr

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.21 02:57:34