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

自定义EarlyStopping回调:val_accuracy连续下降指定步数时终止训练

自定义早停Callback的问题修复

你的代码无法终止训练的核心问题有两个:

  1. 指标获取错误:val_accuracy是epoch级的验证指标,只会在每个epoch完成验证后生成,on_train_batch_end的logs里只有训练集的指标(比如accuracy、loss),所以current始终为None,后续逻辑根本没执行。
  2. 逻辑判断写反:原代码中,当指标变好时(比如准确率上升)反而增加等待计数,这完全违背了“连续下降才停止”的需求。

以下是针对不同需求的修复方案:

方案一:监控训练集batch级准确率(连续N步下降停止)

如果你需要监控每个训练batch的准确率,连续下降指定步数后终止训练:

class CustomEarlyStopping(tf.keras.callbacks.Callback):
    def __init__(self, monitor, max_steps, mode='min', delta=0):
        super().__init__()
        self.monitor = monitor
        self.max_steps = max_steps
        self.mode = mode
        self.delta = delta
        self.wait = 0
        self.stopped_step = 0
        self.best = None
        self.steps_per_epoch = None

    def on_train_begin(self, logs=None):
        self.wait = 0
        self.stopped_step = 0
        self.best = None
        # 获取每个epoch的总步数
        self.steps_per_epoch = self.params['steps']

    def on_train_batch_end(self, batch, logs=None):
        current = logs.get(self.monitor)
        if current is None:
            print(f"警告:监控指标{self.monitor}不在batch日志中")
            return

        if self.best is None:
            self.best = current
        else:
            # 判断当前指标是否变差
            is_worse = False
            if self.mode == 'min':
                # 越小越好的指标(如loss),当前值比最佳值大delta以上=变差
                is_worse = current > self.best + self.delta
            elif self.mode == 'max':
                # 越大越好的指标(如accuracy),当前值比最佳值小delta以上=变差
                is_worse = current < self.best - self.delta

            if is_worse:
                self.wait += 1
                if self.wait >= self.max_steps:
                    self.stopped_step = batch + 1  # batch从0开始计数
                    self.model.stop_training = True
                    print(f"在第{self.stopped_step}步触发早停")
            else:
                # 指标变好,重置等待计数并更新最佳值
                self.wait = 0
                self.best = current

# 使用示例:监控训练集准确率,连续100步下降则停止
early_stopping = CustomEarlyStopping(monitor='accuracy', max_steps=100, mode='max')

方案二:监控验证集准确率(每N步验证,连续多次下降停止)

如果坚持要监控val_accuracy,需要手动在指定步数后触发验证,再判断指标是否连续下降:

class CustomValEarlyStopping(tf.keras.callbacks.Callback):
    def __init__(self, monitor, check_every_steps, max_consecutive_drops, delta=0):
        super().__init__()
        self.monitor = monitor
        self.check_every_steps = check_every_steps  # 每多少步验证一次
        self.max_consecutive_drops = max_consecutive_drops  # 连续下降多少次后停止
        self.delta = delta
        self.wait_drops = 0
        self.best_val = None
        self.total_steps = 0

    def on_train_batch_end(self, batch, logs=None):
        self.total_steps += 1
        # 每check_every_steps步触发一次验证
        if self.total_steps % self.check_every_steps == 0:
            # 手动执行验证(需确保fit时传入了validation_data)
            val_logs = self.model.evaluate(self.validation_data, verbose=0)
            # 从验证结果中提取监控指标
            val_idx = self.model.metrics_names.index(self.monitor)
            current_val = val_logs[val_idx]

            if self.best_val is None:
                self.best_val = current_val
                print(f"第{self.total_steps}步 | 最佳{self.monitor}: {current_val:.4f}")
            else:
                # 判断验证指标是否下降
                is_drop = False
                if 'accuracy' in self.monitor:
                    # 准确率:当前值比最佳值小delta以上=下降
                    is_drop = current_val < self.best_val - self.delta
                else:
                    # 损失类:当前值比最佳值大delta以上=下降
                    is_drop = current_val > self.best_val + self.delta

                if is_drop:
                    self.wait_drops += 1
                    print(f"第{self.total_steps}步 | {self.monitor}下降 | 连续下降次数: {self.wait_drops}")
                    if self.wait_drops >= self.max_consecutive_drops:
                        print(f"连续{self.max_consecutive_drops}次验证{self.monitor}下降,触发早停")
                        self.model.stop_training = True
                else:
                    self.wait_drops = 0
                    self.best_val = current_val
                    print(f"第{self.total_steps}步 | {self.monitor}提升 | 新最佳值: {current_val:.4f}")

# 使用示例:每50步验证一次val_accuracy,连续3次下降则停止
early_stopping = CustomValEarlyStopping(
    monitor='val_accuracy',
    check_every_steps=50,
    max_consecutive_drops=3,
    delta=0.001
)
# 注意:调用model.fit()时必须传入validation_data参数

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 03:09:27