PyTorch如何实现Keras EarlyStopping并按自定义指标保存最优模型
PyTorch自定义EarlyStopping实现方案
完全可以自己手写一个简易的EarlyStopping类,灵活支持任意自定义监控指标,不用依赖第三方库,代码逻辑和Keras的EarlyStopping完全对齐,以下是可直接复用的实现:
第一步:实现EarlyStopping类
import torch import numpy as np class EarlyStopping: def __init__(self, monitor='f1_score', min_delta=0, patience=10, verbose=0, mode='max'): """ 参数和Keras的EarlyStopping完全对应 :param monitor: 要监控的指标名称 :param min_delta: 指标变化小于该值视为没有提升 :param patience: 连续多少轮没有提升就触发早停 :param verbose: 0不输出日志,1输出日志 :param mode: 可选'min'/'max',min是指标越小越好(比如损失),max是指标越大越好(比如f1、准确率) """ self.monitor = monitor self.min_delta = min_delta self.patience = patience self.verbose = verbose self.mode = mode self.best_score = None self.early_stop = False self.counter = 0 self.best_model_state = None # 初始化指标比较逻辑 if self.mode == 'min': self.compare = lambda a, b: a < b - self.min_delta elif self.mode == 'max': self.compare = lambda a, b: a > b + self.min_delta else: raise ValueError("mode只能是'min'或者'max'") def __call__(self, current_metric, model): score = current_metric if self.best_score is None: self.best_score = score self.save_best_model(model) elif not self.compare(score, self.best_score): self.counter += 1 if self.verbose: print(f"EarlyStopping counter: {self.counter}/{self.patience},当前最优{self.monitor}: {self.best_score:.4f}") if self.counter >= self.patience: self.early_stop = True else: self.best_score = score self.save_best_model(model) self.counter = 0 def save_best_model(self, model): self.best_model_state = model.state_dict() if self.verbose: print(f"最优{self.monitor}更新为: {self.best_score:.4f},已保存模型权重")
注意你示例的Keras代码里监控f1_score用了
mode='min'是错误的,f1是越高越好,mode应该设为'max',只有监控损失这种越低越好的指标才用'min'。
第二步:训练流程中接入使用
from sklearn.metrics import f1_score # 初始化早停类,和你需要的Keras代码逻辑对齐,修正了mode为max early_stopping = EarlyStopping(monitor='f1_score', patience=10, verbose=1, mode='max') # 训练循环 for epoch in range(总训练轮数): # 训练步骤 model.train() for batch in 训练集: # 前向传播、反向传播、更新权重代码 # 验证步骤,计算自定义f1_score model.eval() all_preds = [] all_labels = [] with torch.no_grad(): for batch in 验证集: inputs, labels = batch outputs = model(inputs) preds = torch.argmax(outputs, dim=1) all_preds.extend(preds.cpu().numpy()) all_labels.extend(labels.cpu().numpy()) # 计算当前轮验证集f1 current_f1 = f1_score(all_labels, all_preds, average='macro') # 传入早停类判断 early_stopping(current_f1, model) # 触发早停就跳出训练循环 if early_stopping.early_stop: print(f"触发早停,最优f1_score为: {early_stopping.best_score:.4f}") break # 训练结束后加载最优模型权重 model.load_state_dict(early_stopping.best_model_state)
功能说明
- 完全适配任意自定义监控指标,除了f1之外,召回率、精确率、mAP等指标都可以直接用,只需要在验证阶段算出对应指标传给EarlyStopping实例即可
- 自动保存最优轮次的模型权重,不需要手动写保存逻辑
- 参数和Keras的EarlyStopping完全对齐,迁移成本极低
内容的提问来源于stack exchange,提问作者Anđela Todorović
相关产品推荐
相关产品推荐

