news 2026/9/10 20:54:48

Detectron2 训练实战指南:自定义训练循环、Trainer 抽象与 Hook 机制全解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Detectron2 训练实战指南:自定义训练循环、Trainer 抽象与 Hook 机制全解析

Detectron2 训练实战指南:自定义训练循环、Trainer 抽象与 Hook 机制全解析

【免费下载链接】detectron2Detectron2 is a platform for object detection, segmentation and other visual recognition tasks.项目地址: https://gitcode.com/GitHub_Trending/de/detectron2

导读

在完成自定义模型与数据加载器的搭建之后,如何高效地把它们跑起来训练,是每个 Detectron2 使用者都会面对的问题。本文以官方训练教程为核心,系统讲解 Detectron2 提供的两种主流训练风格——自由度极高的"自定义训练循环"(Custom Training Loop)与内置标准行为的 Trainer 抽象(SimpleTrainer / DefaultTrainer),并深入剖析其背后的 Hook 机制与指标日志(EventStorage / EventWriter)实现。读完本文,你将掌握如何从零手写训练循环、如何用几行代码定制 DefaultTrainer、如何编写自定义 Hook 扩展训练行为,以及如何在训练过程中向 TensorBoard 与 JSON 日志写入自定义指标。

两种训练风格的选择

Detectron2 官方训练教程开篇即指出:当你已经拥有一个模型(model)与一个数据加载器(data loader)之后,运行训练通常有两种偏好风格。二者的本质区别在于"框架替你做了多少",以及"你愿意接受多少默认假设"。从工程实践角度看,这一选择直接决定了后续做研究时修改训练逻辑的成本。

风格一:自定义训练循环(Custom Training Loop)

当模型与数据加载器就绪后,编写训练循环所需的其余一切几乎都可以从 PyTorch 原生的 API 中找到,你可以完全自由地手写训练循环。这种风格让研究人员能够更清晰地掌控全部训练逻辑,获得完全的支配权——任何针对训练逻辑的定制都可以直接由用户自己控制,不必绕过框架的抽象层。

仓库中提供了一个完整示例:tools/plain_train_net.py。该脚本的文档字符串明确说明了它的定位:它读取给定配置文件并运行训练或评估,是能够训练 Detectron2 标准模型的入口脚本;但它内含大量针对内置模型的特殊逻辑(例如根据数据集元数据evaluator_type做 if-else 分派评估器的get_evaluator),因此未必适合你自己的项目。官方建议把 Detectron2 当作库来使用,将这个文件作为"如何使用库"的示例,再根据自己的数据集与定制需求编写自己的脚本。它相比train_net.py支持更少的默认特性、包含更少的抽象层,因此更容易加入自定义逻辑。

风格二:Trainer 抽象(Trainer Abstraction)

Detectron2 同时提供了一个标准化的 Trainer 抽象,配合 Hook 系统来简化标准训练行为,内置两种实例化:

  • SimpleTrainer:提供"单损失(single-cost)、单优化器(single-optimizer)、单数据源(single-data-source)"场景下最小化的训练循环,除此之外什么都不做。checkpointing、logging 等其他任务都可以通过 Hook 系统实现。
  • DefaultTrainer:由 yacs 配置初始化的SimpleTrainer,被 tools/train_net.py 及大量脚本使用。它包含了更多人们通常希望默认开启的标准行为,例如优化器、学习率调度、日志、评估、checkpoint 等的默认配置。

从源码看,DefaultTrainer直接继承自TrainerBase(见 detectron2/engine/train_loop.py),其构造函数按固定顺序完成build_modelbuild_optimizerbuild_train_loader,随后根据cfg.SOLVER.AMP.ENABLED选择AMPTrainer还是SimpleTrainer作为底层_trainer,并注册一组默认 Hook(详见 detectron2/engine/defaults.py)。官方在类注释中坦诚提醒:这个类做了"许多假设",任何超出SimpleTrainer的假设对于研究来说都可能过多,一旦不适用,鼓励使用者重写其方法、改用SimpleTrainer,或直接照搬plain_train_net.py写自己的训练循环。

深入剖析自定义训练循环:plain_train_net.py 源码拆解

plain_train_net.pydo_train函数是理解手写训练循环的最佳教材,其核心流程如下(见 tools/plain_train_net.py):

def do_train(cfg, model, resume=False): model.train() optimizer = build_optimizer(cfg, model) scheduler = build_lr_scheduler(cfg, optimizer) checkpointer = DetectionCheckpointer( model, cfg.OUTPUT_DIR, optimizer=optimizer, scheduler=scheduler ) start_iter = ( checkpointer.resume_or_load(cfg.MODEL.WEIGHTS, resume=resume).get("iteration", -1) + 1 ) max_iter = cfg.SOLVER.MAX_ITER periodic_checkpointer = PeriodicCheckpointer( checkpointer, cfg.SOLVER.CHECKPOINT_PERIOD, max_iter=max_iter ) writers = default_writers(cfg.OUTPUT_DIR, max_iter) if comm.is_main_process() else [] data_loader = build_detection_train_loader(cfg) logger.info("Starting training from iteration {}".format(start_iter)) with EventStorage(start_iter) as storage: for data, iteration in zip(data_loader, range(start_iter, max_iter)): storage.iter = iteration loss_dict = model(data) losses = sum(loss_dict.values()) assert torch.isfinite(losses).all(), loss_dict loss_dict_reduced = {k: v.item() for k, v in comm.reduce_dict(loss_dict).items()} losses_reduced = sum(loss for loss in loss_dict_reduced.values()) if comm.is_main_process(): storage.put_scalars(total_loss=losses_reduced, **loss_dict_reduced) optimizer.zero_grad() losses.backward() optimizer.step() storage.put_scalar("lr", optimizer.param_groups[0]["lr"], smoothing_hint=False) scheduler.step() if ( cfg.TEST.EVAL_PERIOD > 0 and (iteration + 1) % cfg.TEST.EVAL_PERIOD == 0 and iteration != max_iter - 1 ): do_test(cfg, model) comm.synchronize() if iteration - start_iter > 5 and ( (iteration + 1) % 20 == 0 or iteration == max_iter - 1 ): for writer in writers: writer.write() periodic_checkpointer.step(iteration)

这段代码几乎浓缩了所有"标准训练"要素,可拆解为以下关键点:

  1. 构建组件:通过build_optimizerbuild_lr_scheduler构造优化器与学习率调度器,通过DetectionCheckpointer加载/恢复权重,并借助PeriodicCheckpointercfg.SOLVER.CHECKPOINT_PERIOD周期保存 checkpoint。
  2. 训练主循环:外层使用zip(data_loader, range(start_iter, max_iter))同时迭代数据与迭代序号,天然支持从断点恢复;model(data)返回一个 loss 字典,sum(loss_dict.values())得到总损失;随后是标准的zero_grad → backward → step三段式,并在每次迭代后执行scheduler.step()
  3. 指标记录:整个循环被包裹在with EventStorage(start_iter) as storage:上下文内,循环中通过storage.put_scalars(...)记录损失、通过storage.put_scalar("lr", ...)记录学习率。
  4. 周期评估与写出:当cfg.TEST.EVAL_PERIOD > 0且到达评估周期时调用do_test;每 20 个迭代(以及最后一个迭代)调用所有default_writerswrite()把指标落盘。

该脚本的另一特色是其基于数据集元数据自动构建评估器的get_evaluator逻辑(见 tools/plain_train_net.py):它会读取MetadataCatalog.get(dataset_name).evaluator_type,据此分派SemSegEvaluatorCOCOEvaluatorCOCOPanopticEvaluatorCityscapesInstanceEvaluatorPascalVOCDetectionEvaluatorLVISEvaluator等,多个评估器通过DatasetEvaluators组合。官方在注释中特意说明:这种 hacky 的 if-else 只是为内置数据集服务,你自己的数据集直接在脚本里手动创建评估器即可。

DefaultTrainer 的定制之道

简单定制:覆写类方法

对于简单定制(例如更换优化器、评估器、LR 调度器、数据加载器等),官方建议像 tools/train_net.py 那样,在子类中覆写DefaultTrainer的对应方法。DefaultTrainer@classmethod形式暴露了如下可覆写入口(见 detectron2/engine/defaults.py):

方法默认实现典型覆写场景
build_model(cfg)调用detectron2.modeling.build_model更换模型架构
build_optimizer(cfg, model)调用detectron2.solver.build_optimizer更换优化器(如 Adam)
build_lr_scheduler(cfg, optimizer)调用detectron2.solver.build_lr_scheduler更换学习率调度策略
build_train_loader(cfg)调用build_detection_train_loader更换数据加载逻辑
build_test_loader(cfg, dataset_name)调用build_detection_test_loader测试时数据预处理定制
build_evaluator(cfg, dataset_name)默认抛NotImplementedError接入自定义数据集评估
build_writers()调用default_writers更换/追加日志写出器

其中build_evaluator默认未实现,会抛出NotImplementedErrortrain_net.pyTrainer(DefaultTrainer)子类覆写了它,按数据集元数据构建对应评估器,同时还额外提供了test_with_TTA方法——当cfg.TEST.AUG.ENABLED开启时,用GeneralizedRCNNWithTTA包装模型做测试时增强(TTA)评估(见 tools/train_net.py)。

train_net.pymain展示了完整的训练入口模式(见 tools/train_net.py):

trainer = Trainer(cfg) trainer.resume_or_load(resume=args.resume) if cfg.TEST.AUG.ENABLED: trainer.register_hooks( [hooks.EvalHook(0, lambda: trainer.test_with_TTA(cfg, trainer.model))] ) return trainer.train()

Hook 系统:扩展训练行为的统一入口

对于训练期间的额外任务,官方建议先检查 Hook 系统是否已经支持。HookBase定义于 detectron2/engine/train_loop.py,其生命周期调用顺序为:

hook.before_train() for iter in range(start_iter, max_iter): hook.before_step() trainer.run_step() hook.after_step() iter += 1 hook.after_train()

HookBase提供五个可覆写方法:before_trainafter_trainbefore_stepafter_backwardafter_step。注意官方在源码中强调两点约定:一是 Hook 方法内部可以通过self.trainer(弱引用代理)访问模型、当前迭代、配置等上下文;二是before_step应当只做可忽略不计的轻量工作,耗时操作应放在after_step中,否则会干扰计时类 Hook 的准确性。

教程给出了一个打印 "hello" 的经典示例:

class HelloHook(HookBase): def after_step(self): if self.trainer.iter % 100 == 0: print(f"Hello at iteration {self.trainer.iter}!")

这个例子的背后原理是:TrainerBase.train()会在每次迭代中依次调用before_step()run_step()after_step(),并把storage.itertrainer.iter保持一致(见 detectron2/engine/train_loop.py)。因此self.trainer.iter就是当前迭代号,% 100 == 0即可实现每 100 次迭代打印一次。

仓库内置了一组开箱即用的 Hook(全部定义于 detectron2/engine/hooks.py):

  • IterationTimer:统计每次迭代耗时,训练结束时输出整体训练速度(Overall training speed: ... s / it)与总训练时间。
  • PeriodicWriter:按周期调用所有EventWriterwrite(),默认周期为 20,最后一个迭代也会写出。
  • PeriodicCheckpointer:按cfg.SOLVER.CHECKPOINT_PERIOD周期保存 checkpoint。
  • BestCheckpointer:基于指定验证指标(如bbox/AP50)保存最优权重,支持mode="max"/"min"
  • LRScheduler:执行内置 LR 调度器并把当前学习率写入 storage。
  • EvalHook:按cfg.TEST.EVAL_PERIOD周期执行评估函数,训练结束也会执行一次;评估结果会被压平后写入 storage。
  • PreciseBN:当cfg.TEST.PRECISE_BN.ENABLED且模型含训练态 BN 层时,用真实统计量(而非 EMA 移动平均)更新 BN 参数。
  • TorchProfiler / AutogradProfiler:性能剖析 Hook,可将 trace 导出为 Chrome tracing JSON 或 TensorBoard 可视化。
  • TorchMemoryStats:周期输出 CUDA 显存占用统计。

DefaultTrainer.build_hooks()(见 detectron2/engine/defaults.py)默认注册了IterationTimerLRScheduler、条件性的PreciseBN、主进程上的PeriodicCheckpointerEvalHookPeriodicWriter,其执行顺序经过精心设计:PreciseBN 在 checkpointer 之前(因为其更新需要被保存)、评估在 checkpoint 之后(若评估失败可用已保存权重调试)、writer 在最后(确保评估指标也能被写出)。

何时应该放弃 Trainer + Hook

教程明确给出了边界:使用 trainer + hook 系统意味着总会有一些非标准行为无法被支持,尤其是在研究中。正因如此,官方刻意将 trainer 与 hook 系统保持最小化而非强大——如果任何需求无法通过该系统实现,直接以plain_train_net.py为起点手动实现自定义训练逻辑反而更简单。

SimpleTrainer.run_step()的标准单步逻辑(见 detectron2/engine/train_loop.py)是理解这一边界的钥匙:它仅做"取数据 → 前向得到 loss 字典 →losses.backward()optimizer.step()",并把zero_grad的位置、梯度累积、梯度裁剪等交由用户通过包装 optimizer 或模型实现。此外,AMPTrainerSimpleTrainer基础上用torch.cuda.amp.autocastGradScaler实现了自动混合精度训练,由cfg.SOLVER.AMP.ENABLED控制切换。

指标日志机制:EventStorage 与 EventWriter

在模型内部写入自定义指标

训练期间,Detectron2 的模型与 trainer 会把指标统一放入一个集中的EventStorage。教程给出的用法如下:

from detectron2.utils.events import get_event_storage # inside the model: if self.training: value = # compute the value from inputs storage = get_event_storage() storage.put_scalar("some_accuracy", value)

其底层实现中,get_event_storage()返回当前上下文栈顶的EventStorage对象,put_scalar(name, value, smoothing_hint=True)会把标量写入以name命名的HistoryBuffer,并附带一个"是否需要平滑"的提示(默认 True,因为多数标量需要平滑才能看出趋势;像学习率这类本身不平滑的信号可传smoothing_hint=False,见 detectron2/utils/events.py)。EventStorage还支持put_scalars(**kwargs)批量写入、put_image向 TensorBoard 添加图像、put_histogram记录直方图,以及name_scope为指标名加前缀便于分组(例如在某个子模块作用域内记录的指标会自动带上模块名/前缀)。

需要特别留意的是调用约束:get_event_storage()必须在with EventStorage(...):上下文内调用,否则会直接断言报错。这正是SimpleTrainerDefaultTrainer以及plain_train_net.py都把整个训练循环包在with EventStorage(start_iter) as storage:里的原因。

指标的多端写出:EventWriter

写入EventStorage的指标随后由各种EventWriter分发到不同目的地。DefaultTrainer默认启用一组EventWriter,其默认配置来自default_writers(output_dir, max_iter)(见 detectron2/engine/defaults.py),共三个:

Writer作用
CommonMetricPrinter向终端打印迭代时间、ETA、显存、全部 loss 与学习率,使用窗口为 20 的中位数平滑
JSONWriter将指标以"每行一个 JSON"的格式追加写入{OUTPUT_DIR}/metrics.json,便于jq等工具解析
TensorboardXWriter将所有标量(以及图像、直方图)写入 TensorBoard 事件文件

如需自定义(例如增加一个远程日志 writer、改变写出频率),build_writers()就是官方预留的覆写点——DefaultTrainerbuild_hooks()PeriodicWriter(self.build_writers(), period=20)正是通过它构建 writer 列表的。

从配置到运行:训练入口实战

命令行通用参数

无论是train_net.py还是plain_train_net.py,都通过default_argument_parser()(见 detectron2/engine/defaults.py)提供统一的命令行接口:

  • --config-file FILE:指定配置文件路径;
  • --resume:尝试从 checkpoint 目录恢复训练;
  • --eval-only:仅执行评估;
  • --num-gpus N:每台机器的 GPU 数;
  • --num-machines N--machine-rank R:多机训练时的机器总数与当前机器序号;
  • --dist-url URL:分布式后端初始化地址(默认tcp://127.0.0.1:<port>,端口由 uid 哈希确定性生成,便于用户发现孤儿进程);
  • opts:命令行覆盖配置,yacs 配置使用空格分隔的PATH.KEY VALUE形式,LazyConfig 使用path.key=value形式。

典型用法(官方 epilog 示例):

# 单机 8 卡训练 python tools/train_net.py --num-gpus 8 --config-file configs/COCO-Detection/faster_rcnn_R_50_FPN_1x.yaml # 命令行覆盖配置项 python tools/train_net.py --config-file config.yaml MODEL.WEIGHTS /path/to/weight.pth SOLVER.BASE_LR 0.001 # 多机训练(两台机器分别执行) python tools/train_net.py --machine-rank 0 --num-machines 2 --dist-url <URL> --config-file config.yaml python tools/train_net.py --machine-rank 1 --num-machines 2 --dist-url <URL> --config-file config.yaml

default_setup会在启动时完成统一的环境初始化:创建输出目录、设置多 rank 日志、打印环境信息与命令行参数、把完整配置备份到输出目录的config.yaml、按SEED设置各 worker 的随机种子等(见 detectron2/engine/defaults.py)。

训练相关的关键配置项

以下配置键直接决定训练行为,均可通过命令行opts覆盖。以 configs/Base-RCNN-FPN.yaml 为例:

DATASETS: TRAIN: ("coco_2017_train",) TEST: ("coco_2017_val",) SOLVER: IMS_PER_BATCH: 16 # 全局 batch size(含所有 GPU) BASE_LR: 0.02 # 基础学习率 STEPS: (60000, 80000) # 学习率衰减的迭代节点 MAX_ITER: 90000 # 总训练迭代数 INPUT: MIN_SIZE_TRAIN: (640, 672, 704, 736, 768, 800) # 训练时随机采样的短边范围

其他常用训练配置还包括:SOLVER.MOMENTUM(默认 0.9)、SOLVER.WEIGHT_DECAYSOLVER.WEIGHT_DECAY_NORMSOLVER.WEIGHT_DECAY_BIASSOLVER.GAMMA(衰减系数)、SOLVER.WARMUP_FACTOR/SOLVER.WARMUP_ITERS/SOLVER.WARMUP_METHOD(预热策略)、SOLVER.CHECKPOINT_PERIOD(checkpoint 周期)、SOLVER.AMP.ENABLED(混合精度开关)、TEST.EVAL_PERIOD(评估周期)与OUTPUT_DIR(输出目录)。这些键在 detectron2/solver/build.py 的build_optimizerbuild_lr_scheduler中被实际消费:例如build_lr_schedulerWARMUP_METHOD存在时会包一层LRMultiplier,其乘子由WARMUP_FACTORWARMUP_ITERS计算得到。

值得注意的是,不同模型族的默认学习率并不相同:RetinaNet 系列在 configs/Base-RetinaNet.yaml 中显式注释BASE_LR: 0.01 # Note that RetinaNet uses a different default learning rate,与 R-CNN 系列的 0.02 形成对比——这提醒我们在迁移配置到新任务时务必核对基线配置的默认值。

多卡训练的自动缩放

DefaultTrainer构造时还会调用auto_scale_workers(cfg, comm.get_world_size())(见 detectron2/engine/defaults.py):当配置中SOLVER.REFERENCE_WORLD_SIZE与当前实际使用的卡数不一致时,会按比例自动缩放配置——IMS_PER_BATCHBASE_LR乘以scaleMAX_ITERWARMUP_ITERSSTEPSEVAL_PERIODCHECKPOINT_PERIOD除以scale,从而保持单卡 batch size 与总训练步数语义不变。例如参考 8 卡配置(IMS_PER_BATCH: 16, BASE_LR: 0.1, MAX_ITER: 5000)在 16 卡上运行时,会自动缩放为IMS_PER_BATCH: 32, BASE_LR: 0.2, MAX_ITER: 2500。这一特性由是否设置REFERENCE_WORLD_SIZE决定是否启用,研究者在扩展实验规模时无需手工重算超参。

总结与决策建议

综合本教程与仓库源码,选择训练方案时可参考以下决策路径:

  1. 标准模型 + 标准流程:直接使用 tools/train_net.py 与DefaultTrainer,零成本获得优化器、LR 调度、日志、评估、checkpoint 的完整默认行为;
  2. 需要替换某个标准组件:继承DefaultTrainer并覆写build_*系列类方法(见 detectron2/engine/defaults.py),改动面最小;
  3. 需要周期性执行额外任务:优先检查内置 Hook 清单(见 detectron2/engine/hooks.py),或仿照HelloHook编写自定义HookBase子类并通过register_hooks注册;
  4. 训练逻辑偏离标准 SGD 流程较远(多损失、多优化器、多数据源、自定义梯度处理):放弃 trainer 抽象,以 tools/plain_train_net.py 为模板手写训练循环,其中do_train已完整示范了"构建组件 → 迭代训练 → 指标记录 → 周期评估与写出 → 周期 checkpoint"的标准骨架;
  5. 无论选哪种风格,指标都统一经EventStorage汇聚、由EventWriter分发,自定义指标时牢记get_event_storage()必须在with EventStorage(...)上下文中调用。

需要强调的是,Detectron2 官方刻意将 trainer 与 hook 系统保持"最小而非强大",这是其设计哲学:框架负责覆盖 80% 的标准场景,剩下的 20% 研究型需求,官方宁可鼓励你绕开抽象直接写循环,也不愿让抽象层变得臃肿而难以调试。理解这一边界,是在 Detectron2 上进行高效训练开发的关键。

【免费下载链接】detectron2Detectron2 is a platform for object detection, segmentation and other visual recognition tasks.项目地址: https://gitcode.com/GitHub_Trending/de/detectron2

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/10 20:52:10

CANN/GE ACL算子描述API

aclmdlCreateAndGetOpDesc 【免费下载链接】ge GE&#xff08;Graph Engine&#xff09;是面向昇腾的图编译器和执行器&#xff0c;提供了计算图优化、多流并行、内存复用和模型下沉等技术手段&#xff0c;加速模型执行效率&#xff0c;减少模型内存占用。 GE 提供对 PyTorch、…

作者头像 李华
网站建设 2026/9/10 20:50:29

量子力学核心概念与应用解析

1. 量子力学基础概述量子力学是现代物理学的重要分支&#xff0c;研究微观粒子运动规律的理论体系。我第一次接触量子力学是在大学物理实验室&#xff0c;当时用双缝实验观察电子干涉现象&#xff0c;那种颠覆经典物理认知的震撼至今难忘。量子力学不仅改变了我们对物质基本组成…

作者头像 李华
网站建设 2026/9/10 20:49:17

CANN/ge图引擎概念原理

概念原理 【免费下载链接】ge GE&#xff08;Graph Engine&#xff09;是面向昇腾的图编译器和执行器&#xff0c;提供了计算图优化、多流并行、内存复用和模型下沉等技术手段&#xff0c;加速模型执行效率&#xff0c;减少模型内存占用。 GE 提供对 PyTorch、TensorFlow 前端的…

作者头像 李华
网站建设 2026/9/10 20:48:09

黑客为什么不攻击微信和支付宝,反而盯着某些平台下手呢?

黑客为什么不攻击微信和支付宝&#xff0c;反而盯着某些平台下手呢&#xff1f; 近期&#xff0c;快手平台疑似遭遇黑客入侵的消息引发行业热议——有网友反馈账号异常登录、个人信息泄露&#xff0c;部分主播甚至出现直播打赏资金被篡改的情况。事件虽未得到官方最终定性&…

作者头像 李华