如何在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
相关产品推荐
相关产品推荐

