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

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ć

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 18:39:03