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

为股票预测Seq2Seq LSTM模型添加Attention并获取权重

在Seq2Seq LSTM模型中添加Attention层并提取注意力权重

一、修改训练脚本:添加Attention层并保存注意力模型

调整Seq2Seq结构,让Encoder返回所有时间步输出,在Decoder中加入Attention层,同时构建单独的模型用于提取注意力权重。

修改后的训练脚本关键部分:

import tensorflow as tf  # 新增导入

# ... 原有的create_sequences函数保持不变 ...

def train(ticker, forecast_horizon, window_size, epochs=25, batch_size=32, training_split = 0.9):
    # ... 原有的数据加载、预处理代码保持不变 ...

    # 定义Seq2Seq模型架构(修改部分)
    # Encoder:返回所有时间步输出和状态
    encoder_inputs = Input(shape=(window_size, X.shape[2]))
    encoder_lstm = LSTM(50, return_sequences=True, return_state=True, dropout=0.2)
    encoder_outputs, state_h, state_c = encoder_lstm(encoder_inputs)
    encoder_states = [state_h, state_c]

    # Decoder:用Encoder状态初始化,加入Attention层
    decoder_inputs = RepeatVector(forecast_horizon)(state_h)
    decoder_lstm = LSTM(50, return_sequences=True, dropout=0.2)
    decoder_outputs = decoder_lstm(decoder_inputs, initial_state=encoder_states)

    # 添加Attention层:计算Decoder每个时间步与Encoder所有时间步的注意力权重
    attention_layer = Attention()
    attention_output = attention_layer([decoder_outputs, encoder_outputs])

    # 合并Attention输出与Decoder输出
    decoder_concat_input = tf.keras.layers.Concatenate(axis=-1)([decoder_outputs, attention_output])

    # 最终输出层
    decoder_dense = TimeDistributed(Dense(1))
    decoder_outputs = decoder_dense(decoder_concat_input)

    # 主预测模型
    model = Model(encoder_inputs, decoder_outputs)
    # 单独构建注意力权重提取模型
    attention_model = Model(encoder_inputs, attention_layer.attention_weights)

    # 编译与训练
    model.compile(optimizer='adam', loss='mean_squared_error')
    model.fit(X_train, y_train, epochs=epochs, batch_size=batch_size, validation_data=(X_test, y_test))

    # 保存模型:同时保存主模型和注意力模型
    model.save(f'../bin/{ticker}/{forecast_horizon}f_{window_size}w_model.keras')
    attention_model.save(f'../bin/{ticker}/{forecast_horizon}f_{window_size}w_attention_model.keras')
    joblib.dump(scaler_features, f'../bin/{ticker}/{forecast_horizon}f_{window_size}w_scaler_features.pkl')
    joblib.dump(scaler_target, f'../bin/{ticker}/{forecast_horizon}f_{window_size}w_scaler_target.pkl')

二、修改预测脚本:加载注意力模型并提取权重

  1. 更新load函数,同时加载注意力模型:
def load(ticker, window_size, forecast_horizon, root_path):
    model_path = os.path.join(root_path, f'models/bin/{ticker}/{forecast_horizon}f_{window_size}w_model.keras')
    attention_model_path = os.path.join(root_path, f'models/bin/{ticker}/{forecast_horizon}f_{window_size}w_attention_model.keras')
    feat_path = os.path.join(root_path, f'models/bin/{ticker}/{forecast_horizon}f_{window_size}w_scaler_features.pkl')
    target_path = os.path.join(root_path, f'models/bin/{ticker}/{forecast_horizon}f_{window_size}w_scaler_target.pkl')

    model = load_model(model_path)
    attention_model = load_model(attention_model_path)
    scaler_features = joblib.load(feat_path)
    scaler_target = joblib.load(target_path)

    return (model, attention_model, scaler_features, scaler_target)
  1. 更新predict函数,提取注意力权重并加入返回结果:
def predict(ticker, window_size, forecast_horizon, start_date, cached_model, root_path):
    model, attention_model, scaler_features, scaler_target = cached_model
    # ... 原有的数据加载、输入序列准备、LIME特征权重计算代码保持不变 ...

    # 提取注意力权重
    attention_weights = attention_model.predict(input_sequence)
    # 原始权重形状:(1, 30, 120) → 转换为(30,120),即每个预测日对应过去120天的权重
    raw_attention_weights = attention_weights[0]
    # 计算所有预测日的平均注意力权重,映射到对应日期
    avg_attention_weights = np.mean(raw_attention_weights, axis=0)
    input_window_dates = dates_full[start_idx - window_size:start_idx]
    date_attention_map = {date.strftime('%Y-%m-%d'): float(weight) for date, weight in zip(input_window_dates, avg_attention_weights)}

    # 修改返回字典,加入注意力权重
    return {
        'ticker': ticker,
        'window_size': window_size,
        'forecast_horizon': forecast_horizon,
        'start_date': start_date,
        'predicted_prices' : predicted_prices,
        'forecast_dates': forecast_dates,
        'dates_full' : dates_full,
        'actual_prices': actual_prices,
        'actual_prices_full': actual_prices_full,        
        'rmse': rmse,
        'rmspe' : rmspe,
        'feature_importance' : feature_importance,
        'raw_attention_weights': raw_attention_weights.tolist(),  # 每个预测日对应120天的权重
        'avg_date_attention': date_attention_map  # 平均后的日期-注意力权重映射
    }

三、关键说明

  • Attention层逻辑:使用Keras内置Attention层实现Bahdanau加性注意力,计算Decoder每个预测时间步对Encoder所有历史时间步的关注度,权重越高说明该历史时间步对当前预测的影响越大。
  • 权重解读:raw_attention_weights是(30,120)的数组,代表30个预测日分别对过去120天的注意力分配;avg_date_attention是所有预测日对历史日期的平均权重,可直接对应到具体日期做分析。
  • 模型保存注意:必须单独保存注意力模型,因为主模型的输出是预测价格,注意力权重属于中间层输出,需要单独构建模型提取。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 13:34:55