无需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
相关产品推荐
相关产品推荐

