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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.26 23:15:55