为股票预测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')
二、修改预测脚本:加载注意力模型并提取权重
- 更新
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)
- 更新
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
相关产品推荐
相关产品推荐

