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

如何在Detectron2模型中添加早停机制?求解过拟合及after_step()使用问题

用Detectron2的after_step()解决过拟合问题

核心逻辑

Detectron2的after_step()是训练循环中每完成一次梯度更新后触发的钩子方法,你可以通过自定义Hook类实现过拟合缓解逻辑(比如早停、动态正则化、梯度裁剪等),再将Hook注册到训练流程中。


1. 自定义Hook类(以早停为例)

早停是缓解过拟合的经典手段——当验证集loss连续多轮无下降时,提前终止训练,避免模型在训练集上过度拟合。代码示例:

from detectron2.engine import HookBase
from detectron2.utils.logger import log_every_n_seconds
import logging
import torch

class EarlyStoppingHook(HookBase):
    def __init__(self, patience=5, min_delta=0.001):
        self.patience = patience  # 连续多少轮无提升就停止
        self.min_delta = min_delta  # loss下降的最小有效阈值
        self.best_loss = float('inf')
        self.counter = 0

    def after_step(self):
        # 仅在验证周期节点触发(需提前配置SOLVER.CHECKPOINT_PERIOD)
        if self.trainer.iter % self.trainer.cfg.SOLVER.CHECKPOINT_PERIOD == 0:
            # 获取当前验证集loss(需确保训练流程中已计算并记录验证loss)
            val_loss = self.trainer.storage.latest().get('val_loss', None)
            if val_loss is None:
                return
            
            log_every_n_seconds(
                logging.INFO,
                f"Current val loss: {val_loss:.4f}, Best val loss: {self.best_loss:.4f}",
                n=10
            )

            # 更新最佳loss或计数
            if val_loss < self.best_loss - self.min_delta:
                self.best_loss = val_loss
                self.counter = 0
            else:
                self.counter += 1
                if self.counter >= self.patience:
                    log_every_n_seconds(logging.INFO, f"Early stopping triggered after {self.trainer.iter} iterations", n=10)
                    self.trainer.stop()

2. 注册Hook到训练器

初始化训练器后,将自定义Hook添加到训练器的钩子列表中即可生效:

from detectron2.engine import DefaultTrainer
from detectron2.config import get_cfg

# 加载你的配置文件
cfg = get_cfg()
cfg.merge_from_file("path/to/your/config.yaml")
cfg.MODEL.WEIGHTS = "path/to/pretrained/model.pth"
# 配置验证周期(比如每1000步验证一次)
cfg.SOLVER.CHECKPOINT_PERIOD = 1000

# 初始化训练器
trainer = DefaultTrainer(cfg)

# 添加早停Hook
early_stop_hook = EarlyStoppingHook(patience=5)
trainer.register_hooks([early_stop_hook])

# 若需调整Hook执行顺序,可使用insert方法(比如放到最前面)
# trainer.hooks.insert(0, early_stop_hook)

# 启动训练
trainer.resume_or_load(resume=False)
trainer.train()

其他可在after_step()中实现的过拟合缓解逻辑

  • 动态权重衰减:随训练进度逐步增大权重衰减系数,增强正则化效果
class DynamicWeightDecayHook(HookBase):
    def __init__(self, start_wd=0.0001, end_wd=0.001, total_iter=10000):
        self.start_wd = start_wd
        self.end_wd = end_wd
        self.total_iter = total_iter

    def after_step(self):
        current_iter = self.trainer.iter
        # 线性递增权重衰减
        wd = self.start_wd + (self.end_wd - self.start_wd) * (current_iter / self.total_iter)
        for param_group in self.trainer.optimizer.param_groups:
            param_group['weight_decay'] = wd
        log_every_n_seconds(logging.INFO, f"Current weight decay: {wd:.6f}", n=20)
  • 梯度裁剪:限制梯度范数,防止梯度爆炸,稳定训练过程
class GradientClipHook(HookBase):
    def __init__(self, max_norm=1.0):
        self.max_norm = max_norm

    def after_step(self):
        torch.nn.utils.clip_grad_norm_(self.trainer.model.parameters(), self.max_norm)

内容的提问来源于stack exchange,提问作者R.K

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.16 15:50:36