
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 AbstractionDetectron2 同时提供了一个标准化的 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_model→build_optimizer→build_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.py的do_train函数是理解手写训练循环的最佳教材其核心流程如下见 tools/plain_train_net.pydef do_train(cfg, model, resumeFalse): model.train() optimizer build_optimizer(cfg, model) scheduler build_lr_scheduler(cfg, optimizer) checkpointer DetectionCheckpointer( model, cfg.OUTPUT_DIR, optimizeroptimizer, schedulerscheduler ) start_iter ( checkpointer.resume_or_load(cfg.MODEL.WEIGHTS, resumeresume).get(iteration, -1) 1 ) max_iter cfg.SOLVER.MAX_ITER periodic_checkpointer PeriodicCheckpointer( checkpointer, cfg.SOLVER.CHECKPOINT_PERIOD, max_itermax_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_losslosses_reduced, **loss_dict_reduced) optimizer.zero_grad() losses.backward() optimizer.step() storage.put_scalar(lr, optimizer.param_groups[0][lr], smoothing_hintFalse) 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)这段代码几乎浓缩了所有标准训练要素可拆解为以下关键点构建组件通过build_optimizer与build_lr_scheduler构造优化器与学习率调度器通过DetectionCheckpointer加载/恢复权重并借助PeriodicCheckpointer按cfg.SOLVER.CHECKPOINT_PERIOD周期保存 checkpoint。训练主循环外层使用zip(data_loader, range(start_iter, max_iter))同时迭代数据与迭代序号天然支持从断点恢复model(data)返回一个 loss 字典sum(loss_dict.values())得到总损失随后是标准的zero_grad → backward → step三段式并在每次迭代后执行scheduler.step()。指标记录整个循环被包裹在with EventStorage(start_iter) as storage:上下文内循环中通过storage.put_scalars(...)记录损失、通过storage.put_scalar(lr, ...)记录学习率。周期评估与写出当cfg.TEST.EVAL_PERIOD 0且到达评估周期时调用do_test每 20 个迭代以及最后一个迭代调用所有default_writers的write()把指标落盘。该脚本的另一特色是其基于数据集元数据自动构建评估器的get_evaluator逻辑见 tools/plain_train_net.py它会读取MetadataCatalog.get(dataset_name).evaluator_type据此分派SemSegEvaluator、COCOEvaluator、COCOPanopticEvaluator、CityscapesInstanceEvaluator、PascalVOCDetectionEvaluator、LVISEvaluator等多个评估器通过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更换优化器如 Adambuild_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默认未实现会抛出NotImplementedError。train_net.py中Trainer(DefaultTrainer)子类覆写了它按数据集元数据构建对应评估器同时还额外提供了test_with_TTA方法——当cfg.TEST.AUG.ENABLED开启时用GeneralizedRCNNWithTTA包装模型做测试时增强TTA评估见 tools/train_net.py。train_net.py的main展示了完整的训练入口模式见 tools/train_net.pytrainer Trainer(cfg) trainer.resume_or_load(resumeargs.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_train、after_train、before_step、after_backward、after_step。注意官方在源码中强调两点约定一是 Hook 方法内部可以通过self.trainer弱引用代理访问模型、当前迭代、配置等上下文二是before_step应当只做可忽略不计的轻量工作耗时操作应放在after_step中否则会干扰计时类 Hook 的准确性。教程给出了一个打印 hello 的经典示例class HelloHook(HookBase): def after_step(self): if self.trainer.iter % 100 0: print(fHello at iteration {self.trainer.iter}!)这个例子的背后原理是TrainerBase.train()会在每次迭代中依次调用before_step()、run_step()、after_step()并把storage.iter与trainer.iter保持一致见 detectron2/engine/train_loop.py。因此self.trainer.iter就是当前迭代号% 100 0即可实现每 100 次迭代打印一次。仓库内置了一组开箱即用的 Hook全部定义于 detectron2/engine/hooks.pyIterationTimer统计每次迭代耗时训练结束时输出整体训练速度Overall training speed: ... s / it与总训练时间。PeriodicWriter按周期调用所有EventWriter的write()默认周期为 20最后一个迭代也会写出。PeriodicCheckpointer按cfg.SOLVER.CHECKPOINT_PERIOD周期保存 checkpoint。BestCheckpointer基于指定验证指标如bbox/AP50保存最优权重支持modemax/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默认注册了IterationTimer、LRScheduler、条件性的PreciseBN、主进程上的PeriodicCheckpointer、EvalHook与PeriodicWriter其执行顺序经过精心设计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 或模型实现。此外AMPTrainer在SimpleTrainer基础上用torch.cuda.amp.autocast与GradScaler实现了自动混合精度训练由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_hintTrue)会把标量写入以name命名的HistoryBuffer并附带一个是否需要平滑的提示默认 True因为多数标量需要平滑才能看出趋势像学习率这类本身不平滑的信号可传smoothing_hintFalse见 detectron2/utils/events.py。EventStorage还支持put_scalars(**kwargs)批量写入、put_image向 TensorBoard 添加图像、put_histogram记录直方图以及name_scope为指标名加前缀便于分组例如在某个子模块作用域内记录的指标会自动带上模块名/前缀。需要特别留意的是调用约束get_event_storage()必须在with EventStorage(...):上下文内调用否则会直接断言报错。这正是SimpleTrainer、DefaultTrainer以及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()就是官方预留的覆写点——DefaultTrainer的build_hooks()中PeriodicWriter(self.build_writers(), period20)正是通过它构建 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.keyvalue形式。典型用法官方 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.yamldefault_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_DECAY、SOLVER.WEIGHT_DECAY_NORM、SOLVER.WEIGHT_DECAY_BIAS、SOLVER.GAMMA衰减系数、SOLVER.WARMUP_FACTOR/SOLVER.WARMUP_ITERS/SOLVER.WARMUP_METHOD预热策略、SOLVER.CHECKPOINT_PERIODcheckpoint 周期、SOLVER.AMP.ENABLED混合精度开关、TEST.EVAL_PERIOD评估周期与OUTPUT_DIR输出目录。这些键在 detectron2/solver/build.py 的build_optimizer与build_lr_scheduler中被实际消费例如build_lr_scheduler在WARMUP_METHOD存在时会包一层LRMultiplier其乘子由WARMUP_FACTOR与WARMUP_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_BATCH与BASE_LR乘以scaleMAX_ITER、WARMUP_ITERS、STEPS、EVAL_PERIOD、CHECKPOINT_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决定是否启用研究者在扩展实验规模时无需手工重算超参。总结与决策建议综合本教程与仓库源码选择训练方案时可参考以下决策路径标准模型 标准流程直接使用 tools/train_net.py 与DefaultTrainer零成本获得优化器、LR 调度、日志、评估、checkpoint 的完整默认行为需要替换某个标准组件继承DefaultTrainer并覆写build_*系列类方法见 detectron2/engine/defaults.py改动面最小需要周期性执行额外任务优先检查内置 Hook 清单见 detectron2/engine/hooks.py或仿照HelloHook编写自定义HookBase子类并通过register_hooks注册训练逻辑偏离标准 SGD 流程较远多损失、多优化器、多数据源、自定义梯度处理放弃 trainer 抽象以 tools/plain_train_net.py 为模板手写训练循环其中do_train已完整示范了构建组件 → 迭代训练 → 指标记录 → 周期评估与写出 → 周期 checkpoint的标准骨架无论选哪种风格指标都统一经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),仅供参考