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

无需Torchsample在PyTorch中实现Early Stopping的方法咨询

PyTorch无第三方依赖早停实现方案

PyTorch本身没有内置早停接口,你不需要依赖torchsample这类外部库,自己手写一个极简的早停逻辑即可,自定义程度更高,适配你的回归任务也更灵活。


1. 先实现早停工具类

把下面的类复制到你的代码中即可:

import torch
import numpy as np

class EarlyStopping:
    def __init__(self, patience=5, min_delta=0, save_path='best_model.pth'):
        """
        参数说明:
        patience: 容忍验证损失不下降的最大epoch数,超过就触发早停
        min_delta: 验证损失下降超过这个值才算有效提升,避免微小波动误判
        save_path: 最优模型权重的保存路径
        """
        self.patience = patience
        self.min_delta = min_delta
        self.save_path = save_path
        self.counter = 0
        self.best_loss = np.inf
        self.early_stop = False

    def __call__(self, val_loss, model):
        if val_loss < self.best_loss - self.min_delta:
            # 验证损失下降,更新最优损失和模型权重
            self.best_loss = val_loss
            torch.save(model.state_dict(), self.save_path)
            self.counter = 0
        else:
            # 验证损失没有有效下降,计数器加1
            self.counter += 1
            if self.counter >= self.patience:
                self.early_stop = True

2. 修正现有代码的小问题

你的现有训练代码存在缩进错误:验证循环被放在了训练batch循环内部,导致每训练一个batch就跑一次全量验证,训练效率会非常低,需要把验证逻辑移到训练batch循环外部,每个epoch跑完所有训练batch再跑验证。


3. 集成早停到训练流程

修改你的训练部分代码如下:

# 初始化早停实例,这里patience设为5可根据需求调整
early_stopping = EarlyStopping(patience=5, min_delta=1e-6, save_path='best_regression_model.pth')

print("Begin Training")
for e in tqdm(range(1, EPOCHS+1)):
    # 训练阶段
    train_epoch_loss = 0
    model.train()
    for X_train_batch, y_train_batch in train_loader:
        X_train_batch, y_train_batch = X_train_batch.to(device), y_train_batch.to(device)
        optimizer.zero_grad()
        
        y_train_pred = model(X_train_batch)
        train_loss = criterion(y_train_pred, y_train_batch.unsqueeze(1))
        train_loss.backward()
        optimizer.step()
        
        train_epoch_loss += train_loss.item()
    
    # 验证阶段(移到训练batch循环外)
    val_epoch_loss = 0
    model.eval()
    with torch.no_grad():
        for X_val_batch, y_val_batch in val_loader:
            X_val_batch, y_val_batch = X_val_batch.to(device), y_val_batch.to(device)
            y_val_pred = model(X_val_batch)
            val_loss = criterion(y_val_pred, y_val_batch.unsqueeze(1))
            val_epoch_loss += val_loss.item()
    
    # 计算平均损失
    avg_train_loss = train_epoch_loss / len(train_loader)
    avg_val_loss = val_epoch_loss / len(val_loader)
    loss_stats["train"].append(avg_train_loss)
    loss_stats["val"].append(avg_val_loss)
    print(f"Epoch {e}: Train loss: {avg_train_loss:.5f} Val loss: {avg_val_loss:.5f}")
    
    # 调用早停逻辑
    early_stopping(avg_val_loss, model)
    if early_stopping.early_stop:
        print(f"验证损失连续{early_stopping.patience}个epoch未提升,触发早停")
        break

# 训练结束后加载最优模型权重
model.load_state_dict(torch.load('best_regression_model.pth'))

后续测试、可视化的代码不需要修改,直接用你原来的即可。如果你想跟踪验证集R²而不是损失作为早停判断指标,只需要修改EarlyStopping类的判断逻辑即可,灵活度很高。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 06:27:01