如何修改Python LSTM模型代码实现年度预测而非月度预测
年度时间序列LSTM预测代码修改方案
针对你的年度数据集,以下是适配年度预测的代码修改,核心是完善未来一年的预测逻辑,并根据年度数据特性调整训练参数:
1. 数据预处理与序列构建(修改后)
from sklearn.preprocessing import MinMaxScaler import numpy as np # Normalize the data scaler = MinMaxScaler() scaled_data = scaler.fit_transform(df_final) # 定义序列长度:建议用过去N年的数据预测下一年,比如设为3(可根据数据量调整) sequence_length = 3 num_features = len(df_final.columns) # Create sequences and corresponding labels sequences = [] labels = [] # 遍历范围调整为len(scaled_data) - sequence_length,确保能取到下一年的标签 for i in range(len(scaled_data) - sequence_length): # 取连续sequence_length年的所有特征作为输入序列 seq = scaled_data[i:i+sequence_length, :] # 目标是下一年的'Indnpa'(索引11) label = scaled_data[i+sequence_length, 11] sequences.append(seq) labels.append(label) # Convert to numpy arrays sequences = np.array(sequences) labels = np.array(labels) # Split into train and test sets train_size = int(0.8 * len(sequences)) train_x, test_x = sequences[:train_size], sequences[train_size:] train_y, test_y = labels[:train_size], labels[train_size:] print("Train X shape:", train_x.shape) print("Train Y shape:", train_y.shape) print("Test X shape:", test_x.shape) print("Test Y shape:", test_y.shape)
说明:调整sequence_length为更合理的数值(如3),让模型学习多年的时序规律,年度数据量通常较少,避免设置过大导致样本不足。
2. 模型构建(基本保留,适配输入形状)
from tensorflow.keras.models import Sequential from tensorflow.keras.layers import LSTM, Dense, Dropout from tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint # Create the LSTM model model = Sequential() # 输入形状由train_x自动推导,无需手动修改 model.add(LSTM(units=128, input_shape=(train_x.shape[1], train_x.shape[2]), return_sequences=True)) model.add(Dropout(0.2)) model.add(LSTM(units=64, return_sequences=True)) model.add(Dropout(0.2)) model.add(LSTM(units=32, return_sequences=False)) model.add(Dropout(0.2)) # 输出层保持1个单元,对应单值预测 model.add(Dense(units=1)) # Compile the model model.compile(optimizer='adam', loss='mean_squared_error')
3. 模型训练(适配年度数据调整参数)
early_stopping = EarlyStopping(monitor='val_loss', patience=5, restore_best_weights=True) model_checkpoint = ModelCheckpoint('/content/drive/MyDrive/dl/weather_prediction/best_model_weights.h5', monitor='val_loss', save_best_only=True) # 年度数据样本量少,调小batch_size避免训练不稳定 history = model.fit( train_x, train_y, epochs=100, batch_size=8, # 从64改为8/16 validation_split=0.2, callbacks=[early_stopping, model_checkpoint] )
说明:年度数据通常样本数远少于月度,过大的batch_size会导致模型无法有效学习,调小后训练更稳定。
4. 下一年度预测实现(新增核心逻辑)
# 加载最优权重 model.load_weights('/content/drive/MyDrive/dl/weather_prediction/best_model_weights.h5') # 获取最新的sequence_length年数据作为输入 latest_sequence = scaled_data[-sequence_length:, :] # 调整形状为模型输入要求:(1, sequence_length, num_features) latest_sequence = latest_sequence.reshape(1, sequence_length, num_features) # 预测下一年的归一化值 predicted_scaled = model.predict(latest_sequence) # 反归一化得到真实NPA值 # 注意:需要构造与scaler输入形状匹配的数组,仅目标列对应预测值,其他列补0不影响反归一化 predicted_reshaped = np.zeros((1, num_features)) predicted_reshaped[0, 11] = predicted_scaled[0, 0] next_year_pred = scaler.inverse_transform(predicted_reshaped)[0, 11] print(f"下一年度NPA预测值:{next_year_pred:.4f}")
5. 可视化修改(包含预测结果)
import matplotlib.pyplot as plt # 测试集预测(用于对比) test_pred_scaled = model.predict(test_x) # 反归一化测试集预测值 test_pred_reshaped = np.zeros((len(test_pred_scaled), num_features)) test_pred_reshaped[:, 11] = test_pred_scaled[:, 0] test_pred = scaler.inverse_transform(test_pred_reshaped)[:, 11] # 真实测试集值 test_true_reshaped = np.zeros((len(test_y), num_features)) test_true_reshaped[:, 11] = test_y test_true = scaler.inverse_transform(test_true_reshaped)[:, 11] # 绘图:包含历史真实值、测试集预测值、下一年预测值 plt.figure(figsize=(12, 6)) # 绘制全部历史真实值 plt.plot(df_final.index, df_final.iloc[:, 11], label='历史实际值') # 绘制测试集预测值(对应测试集时间范围) test_index = df_final.index[train_size + sequence_length:] plt.plot(test_index, test_pred, label='测试集预测值') # 添加下一年预测点 next_year_index = df_final.index[-1] + 1 # 假设索引是年份整数,根据实际索引类型调整 plt.scatter(next_year_index, next_year_pred, color='red', label='下一年预测值', s=100) plt.title('行业NPA预测对比') plt.xlabel('年份') plt.ylabel('NPA比率') plt.legend() plt.show()
说明:根据你的df_final.index类型调整next_year_index,如果是datetime索引,用df_final.index[-1] + pd.DateOffset(years=1)。
内容的提问来源于stack exchange,提问作者Malinda Liyanage
相关产品推荐
相关产品推荐

