自定义EarlyStopping回调:val_accuracy连续下降指定步数时终止训练
自定义早停Callback的问题修复
你的代码无法终止训练的核心问题有两个:
- 指标获取错误:
val_accuracy是epoch级的验证指标,只会在每个epoch完成验证后生成,on_train_batch_end的logs里只有训练集的指标(比如accuracy、loss),所以current始终为None,后续逻辑根本没执行。 - 逻辑判断写反:原代码中,当指标变好时(比如准确率上升)反而增加等待计数,这完全违背了“连续下降才停止”的需求。
以下是针对不同需求的修复方案:
方案一:监控训练集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
相关产品推荐
相关产品推荐

