PyTorch Lightning LSTM模型生产预测及结果重复问题求助
LSTM回归模型预测问题修复与生产环境指南
一、预测结果全相同的问题修复
1. 维护LSTM的隐藏状态
原模型forward方法每次调用都会重置LSTM的隐藏状态,导致多步预测时模型无法利用历史序列的上下文信息,最终输出趋于一致。修改模型的forward方法,支持传递和更新隐藏状态:
class LSTMRegressor(L.LightningModule): # ... 其他代码保持不变 ... def forward(self, x, hidden=None): # 首次调用时初始化隐藏状态 if hidden is None: batch_size = x.size(0) hidden = ( torch.zeros(self.num_layers, batch_size, self.hidden_size, device=x.device), torch.zeros(self.num_layers, batch_size, self.hidden_size, device=x.device) ) lstm_out, hidden = self.lstm(x, hidden) y_pred = self.fc(lstm_out[:, -1]) return y_pred, hidden
在预测脚本中维护隐藏状态,确保每一步预测都基于上下文更新:
# ... 模型加载部分 ... model.eval() # 初始化输入序列(匹配模型训练时的seq_len) new_data = np.array([[5.45], [5.43], [5.45], [5.43], [5.36], [5.33], [5.21]]) new_data_tensor = torch.tensor(new_data, dtype=torch.float32).unsqueeze(0) # 形状: (1, seq_len, n_features) hidden_state = None predictions = [] with torch.no_grad(): for i in range(14): prediction, hidden_state = model(new_data_tensor, hidden_state) print(f"第{i+1}步预测值: {prediction}") # 调整预测值维度,匹配输入序列格式 pred_reshaped = prediction.unsqueeze(1) # 更新输入序列:移除最旧的时间步,添加新预测值 new_data_tensor = torch.cat((new_data_tensor[:, 1:, :], pred_reshaped), dim=1) predictions.append(prediction.cpu().numpy())
2. 修复模型加载参数错误
原load_from_checkpoint中参数传递错误,传入的是字符串列表而非实际数值,需从config的p对象中读取真实参数:
model = LSTMRegressor.load_from_checkpoint( "./checkpoints/model-epoch=29-val_loss=11.96.ckpt", n_features=p.n_features, hidden_size=p.hidden_size, criterion=p.criterion, num_layers=p.num_layers, dropout=p.dropout, learning_rate=p.learning_rate, output_size=p.output_size, )
同时修正train.py中的模型和DataModule初始化代码,传入实际参数值:
# train.py中DataModule初始化 dm = EpiCountsDataModule( seq_len=p.seq_len, batch_size=p.batch_size, num_workers=p.num_workers, ) # 模型初始化 model = LSTMRegressor( n_features=p.n_features, hidden_size=p.hidden_size, criterion=p.criterion, num_layers=p.num_layers, dropout=p.dropout, learning_rate=p.learning_rate, output_size=p.output_size, )
3. 对齐数据预处理逻辑
如果训练时对数据做了归一化/标准化(如MinMaxScaler、StandardScaler),预测时必须使用相同的缩放器处理输入数据,预测后再反缩放得到真实值,否则会导致模型输出异常:
import joblib # 加载训练时保存的缩放器 scaler = joblib.load("scaler.pkl") # 预测前缩放输入 new_data_scaled = scaler.transform(new_data) new_data_tensor = torch.tensor(new_data_scaled, dtype=torch.float32).unsqueeze(0) # 预测后反缩放得到真实值 prediction_np = prediction.cpu().numpy() prediction_original = scaler.inverse_transform(prediction_np)
二、生产环境预测部署指南
1. 模型导出为TorchScript
将PyTorch Lightning模型导出为TorchScript格式,便于生产环境快速加载和推理:
# 加载训练好的模型 model = LSTMRegressor.load_from_checkpoint("your_checkpoint.ckpt", **params) model.eval() # 导出为TorchScript scripted_model = torch.jit.script(model) scripted_model.save("lstm_regressor.pt") # 生产环境加载模型 loaded_model = torch.jit.load("lstm_regressor.pt") loaded_model.eval()
2. 构建推理API
用FastAPI搭建轻量级HTTP API,支持外部系统调用预测:
from fastapi import FastAPI import torch import numpy as np app = FastAPI() # 加载TorchScript模型 model = torch.jit.load("lstm_regressor.pt") model.eval() @app.post("/predict") def predict(sequence: list[list[float]], steps: int = 14): # 转换输入为模型要求的张量格式 seq_tensor = torch.tensor(sequence, dtype=torch.float32).unsqueeze(0) hidden_state = None predictions = [] with torch.no_grad(): for _ in range(steps): pred, hidden_state = model(seq_tensor, hidden_state) predictions.append(pred.cpu().numpy().tolist()) # 更新输入序列 pred_reshaped = pred.unsqueeze(1) seq_tensor = torch.cat((seq_tensor[:, 1:, :], pred_reshaped), dim=1) return {"predictions": predictions}
启动API服务:
uvicorn main:app --host 0.0.0.0 --port 8000
3. 批量预测优化
如果需要处理大量序列,将输入整理成(batch_size, seq_len, n_features)的批量格式,利用GPU并行加速推理,提升处理效率。
4. 监控与维护
- 记录每次预测的输入、输出和耗时,便于排查异常请求。
- 定期用最新数据验证模型性能,当数据分布发生偏移时,重新训练并更新模型。
内容的提问来源于stack exchange,提问作者Ankush
相关产品推荐
相关产品推荐

