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

PyTorch中多变量时间序列LSTM模型的验证损失与早停实现

LSTM油价预测:验证集损失计算与早停机制实现

一、添加验证集损失计算

你当前的训练循环已包含测试集损失计算,只需在eval模式下加入验证集的前向传播与损失计算逻辑即可,具体修改如下:

  1. 在验证/测试阶段加入torch.no_grad()上下文管理器,关闭梯度计算以节省资源
  2. 新增验证集的预测生成与损失计算步骤
  3. 打印日志时加入验证损失,方便监控模型泛化能力

修改后的training_loop函数:

def training_loop(n_epochs, lstm, optimiser, loss_fn, X_train, y_train, X_test, y_test,
                  X_val , y_val, device):
    for epoch in range(n_epochs):
        # 训练阶段
        lstm.train()
        outputs = lstm(X_train)
        optimiser.zero_grad()
        train_loss = loss_fn(outputs, y_train)
        train_loss.backward()
        optimiser.step()
        
        # 验证与测试阶段(关闭梯度计算)
        lstm.eval()
        with torch.no_grad():
            # 计算验证损失
            val_preds = lstm(X_val)
            val_loss = loss_fn(val_preds, y_val)
            # 计算测试损失
            test_preds = lstm(X_test)
            test_loss = loss_fn(test_preds, y_test)
        
        # 每100轮打印训练日志
        if epoch % 100 == 0:
            print("Epoch: %d, train loss: %1.5f, val loss: %1.5f, test loss: %1.5f" % 
                  (epoch, train_loss.item(), val_loss.item(), test_loss.item()))

二、实现早停机制

早停机制通过监控验证损失,当验证损失连续多轮未下降时提前终止训练,避免模型过拟合。核心逻辑为跟踪最佳验证损失、设置耐心值(允许损失不下降的最大轮数),具体实现如下:

  1. 初始化最佳验证损失(设为无穷大)、耐心值、计数器
  2. 每轮验证后,若当前验证损失低于最佳值,则更新最佳损失并保存模型;否则计数器加1
  3. 当计数器达到耐心值时,终止训练并加载最佳模型

修改后的完整训练循环代码:

def training_loop(n_epochs, lstm, optimiser, loss_fn, X_train, y_train, X_test, y_test,
                  X_val , y_val, device, patience=100):
    best_val_loss = float('inf')
    counter = 0
    best_model_path = 'best_lstm_oil_price.pth'  # 最佳模型保存路径
    
    for epoch in range(n_epochs):
        # 训练阶段
        lstm.train()
        outputs = lstm(X_train)
        optimiser.zero_grad()
        train_loss = loss_fn(outputs, y_train)
        train_loss.backward()
        optimiser.step()
        
        # 验证与测试阶段
        lstm.eval()
        with torch.no_grad():
            val_preds = lstm(X_val)
            val_loss = loss_fn(val_preds, y_val)
            test_preds = lstm(X_test)
            test_loss = loss_fn(test_preds, y_test)
        
        # 早停逻辑
        if val_loss < best_val_loss:
            best_val_loss = val_loss
            torch.save(lstm.state_dict(), best_model_path)  # 保存最佳模型参数
            counter = 0  # 重置计数器
        else:
            counter += 1
            if counter >= patience:
                print(f"早停触发!第{epoch}轮验证损失未下降,最佳模型已保存至{best_model_path}")
                break
        
        # 打印训练日志
        if epoch % 100 == 0:
            print("Epoch: %d, train loss: %1.5f, val loss: %1.5f, test loss: %1.5f" % 
                  (epoch, train_loss.item(), val_loss.item(), test_loss.item()))
    
    # 训练结束后加载最佳模型
    lstm.load_state_dict(torch.load(best_model_path))
    print("训练完成,已加载最佳模型")

额外优化建议

  • 移除模型中过时的Variable:PyTorch 0.4.0版本后无需手动封装Variable,同时建议将模型和数据移至同一设备(CPU/GPU):
    # 修改LSTM模型的forward函数
    def forward(self,x):
        h_0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).to(x.device)
        c_0 = torch.zeros(self.num_layers, x.size(0), self.hidden_size).to(x.device)
        output, (hn, cn) = self.lstm(x, (h_0, c_0))
        hn = hn.view(-1, self.hidden_size)
        out = self.relu(hn)
        out = self.fc_1(out)
        out = self.relu(out)
        out = self.fc_2(out)
        return out
    
  • 设备配置:在模型调用前统一设置设备,并将所有数据张量移至该设备:
    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
    # 数据张量移至设备
    X_train_tensors = X_train_tensors.to(device)
    y_train_tensors = y_train_tensors.to(device)
    X_val_tensors = X_val_tensors.to(device)
    y_val_tensors = y_val_tensors.to(device)
    X_test_tensors = X_test_tensors.to(device)
    y_test_tensors = y_test_tensors.to(device)
    # 模型移至设备
    lstm = LSTM(num_classes, input_size, hidden_size, num_layers).to(device)
    
  • 调整耐心值:根据训练情况调整patience参数(如50或150),平衡训练效率与模型效果

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.21 02:35:07