资讯动态

Ultralytics 引擎深入定制指南:基于 BaseTrainer 与 DetectionTrainer 的自定义训练

发布时间:2026/9/10 14:15:55 来源:尧图企业网站定制
Ultralytics 引擎深入定制指南基于 BaseTrainer 与 DetectionTrainer 的自定义训练【免费下载链接】ultralyticsUltralytics YOLO26, YOLO11, YOLOv8 — object detection, instance segmentation, semantic segmentation, image classification, pose estimation, object tracking项目地址: https://gitcode.com/GitHub_Trending/ul/ultralyticsUltralytics 的命令行与 Python 接口本质上是构建于底层引擎执行器之上的一层高层抽象。本文以 docs/en/usage/engine.md 为主体脉络系统讲解BaseTrainer与DetectionTrainer的架构、自定义训练器的标准写法以及如何通过YOLO模型类在保持简洁接口的同时注入自定义训练逻辑并结合仓库源码剖析其背后的实现原理。引言为什么需要理解 Trainer 引擎Ultralytics YOLO 提供了两个面向用户的高层入口——CLI如yolo train ...与 Python 接口如YOLO(yolo26n.pt).train(...)。它们都封装了大量底层逻辑数据集构建、数据加载、批预处理、损失计算、优化器与调度器、验证、模型保存、日志与回调触发等。这些逻辑被统一收拢在一套“引擎执行器”engine executors中其中负责训练的是Trainer家族。从源码结构看这套体系位于 ultralytics/engine/trainer.pyBaseTrainer第 72 行起以及 ultralytics/models/yolo/detect/train.pyDetectionTrainer等按任务划分的子类中。理解 Trainer 引擎是进行“Advanced Customization高级定制”的前提无论是接入自定义模型、自定义数据加载器还是修改损失函数、在特定时机注入回调本质上都是继承对应 Trainer 类并覆写其中若干方法。本文面向的读者是希望为特定任务定制训练流程、想读懂源码中训练循环关键钩子、以及需要将自定义训练器与YOLO高层接口组合使用的开发者。BaseTrainer通用训练例程的基座BaseTrainer提供了一套与具体任务解耦的通用训练例程。它本身不关心“目标检测还是实例分割”而是把差异化环节抽象成可覆写的钩子方法由各任务子类DetectionTrainer、SegmentationTrainer、PoseTrainer、ClassificationTrainer等填充具体实现。核心可覆写方法文档明确给出两个最具代表性的覆写入口其余常用钩子汇总如下方法作用参考来源get_model(cfg, weights)构建待训练的模型文档正文 BaseTrainer 参考get_dataloader()构建训练/验证数据加载器同上preprocess_batch()在模型前向前处理批数据缩放、归一化等同上set_model_attributes()依据数据集信息设置模型属性如类别数、类名同上get_validator()返回用于评估模型的验证器同上在DetectionTrainer的源码 ultralytics/models/yolo/detect/train.py 中可以逐一印证这些钩子的真实形态get_model()第 188 行通过DetectionModel(cfg, nc..., ch...)构建模型并按需model.load(weights)载入预训练权重get_dataloader()第 78 行调用build_dataset()后用build_dataloader()封装并处理rect模式与 shuffle 冲突等细节preprocess_batch()第 107 行将图像张量统一float() / 255归一化并支持multi_scale多尺度随机采样set_model_attributes()第 139 行把数据集类别数、类别名与超参数挂载到模型上。因此如果你的任务不属于任何现成类型最直接的做法就是继承BaseTrainer按约定格式覆写上述函数接入你自己的模型与数据加载器即可。BaseTrainer 初始化时做了什么理解初始化逻辑有助于判断何时、在哪里插入自定义逻辑。阅读 ultralytics/engine/trainer.py 第 123 行的__init__可以发现几个关键事实训练配置通过get_cfg(cfg, overrides)解析合并overrides字典即可覆盖默认参数配置定义见 ultralytics/cfg/default.yaml设备解析、随机种子初始化、保存目录创建与args.yaml落盘在此完成self.last与self.best分别指向保存目录下weights/last.pt与weights/best.pt这也是后续trainer.best属性的来源epochs缺省时回退为 100self.args.epochs or 100回调系统在此提前初始化self.callbacks _callbacks or callbacks.get_default_callbacks()保证on_pretrain_routine_start能尽早捕获原始参数。DetectionTrainer目标检测任务的落地实现DetectionTrainer继承自BaseTrainer补齐了目标检测训练所需的全部细节。其最直接的使用方式如下完整代码见原文档from ultralytics.models.yolo.detect import DetectionTrainer trainer DetectionTrainer(overrides{...}) trainer.train() trained_model trainer.best # Get the best model其中overrides是覆盖默认配置的字典例如dict(modelyolo26n.pt, datacoco8.yaml, epochs3)这与DetectionTrainer类文档字符串中的示例ultralytics/models/yolo/detect/train.py 第 47-52 行一致。coco8.yaml是仓库自带的微型数据集定义ultralytics/cfg/datasets/coco8.yaml适合快速验证训练流程。trainer.train()是训练主入口它会按 ultralytics/engine/trainer.py 中_setup_train()第 315 行→_do_train()的流程依次完成模型初始化、通道内存布局channels_last、set_model_attributes()、层冻结、AMP 自动混精度检测、优化器与调度器装配等步骤再进入逐 epoch 的训练循环。训练完成后trainer.best即指向最佳权重best.pt由验证适应度自动挑选可以直接用于后续推理或导出。覆写 get_model 接入自定义检测模型当默认的DetectionTrainer无法直接支持你的自定义检测模型时可以继承它并只覆写get_modelfrom ultralytics.models.yolo.detect import DetectionTrainer class CustomTrainer(DetectionTrainer): def get_model(self, cfgNone, weightsNone, verboseTrue): Loads a custom detection model given configuration and weight files. trainer CustomTrainer(overrides{...}) trainer.train()由于子类完整继承了父类全部能力你只需聚焦于差异点——这正是“模组化”设计的体现各组件可以独立修改而不破坏整条流水线。同时定制损失函数与回调完整示例原文档给出了一个同时定制“损失函数”与“周期性回调”的进阶示例这里完整保留from ultralytics.models.yolo.detect import DetectionTrainer from ultralytics.nn.tasks import DetectionModel class MyCustomModel(DetectionModel): def init_criterion(self): Initializes the loss function and adds a callback for uploading the model to Google Drive every 10 epochs. class CustomTrainer(DetectionTrainer): def get_model(self, cfgNone, weightsNone, verboseTrue): Returns a customized detection model instance configured with specified config and weights. return MyCustomModel(...) # Callback to upload model weights def log_model(trainer): Logs the path of the last model weight used by the trainer. last_weight_path trainer.last print(last_weight_path) trainer CustomTrainer(overrides{...}) trainer.add_callback(on_train_epoch_end, log_model) # Adds to existing callbacks trainer.train()该示例的技术要点在于DetectionModel位于 ultralytics/nn/tasks.py通过覆写init_criterion()即可替换损失初始化逻辑trainer.last属性对应上一节提到的weights/last.pt这里用它拿到“每个 epoch 结束时的最新权重路径”trainer.add_callback(on_train_epoch_end, log_model)会把回调追加到该事件已有的回调列表末尾对应 ultralytics/engine/trainer.py 第 205 行的self.callbacks[event].append(callback)因此不会覆盖内置回调。关于回调可触发的事件与入口的完整清单参见 Callbacks 指南。扩展阅读BaseTrainer 中还值得覆写的钩子原文档把“用 YOLO 搭配自定义 Trainer”放在较高层级介绍但在真实的定制实践中以下来自 BaseTrainer 及其配套指南的钩子同样常用均有源码佐证validate()自定义指标与“按自定义指标保存模型”默认情况下best.pt由适应度fitness决定检测任务默认取mAP0.5:0.95。若希望按mAP0.5、召回率或自定义 F1 等指标挑选最优模型可覆写validate()让其在调用super().validate()后返回目标指标作为 fitnessimport numpy as np from ultralytics import YOLO from ultralytics.models.yolo.detect import DetectionTrainer from ultralytics.utils import LOGGER class MetricsTrainer(DetectionTrainer): def validate(self): Run validation and compute per-class F1 scores. metrics, fitness super().validate() if metrics is None: return metrics, fitness if hasattr(self.validator, metrics) and hasattr(self.validator.metrics, box): box self.validator.metrics.box f1_per_class box.f1 class_indices box.ap_class_index names self.validator.names mean_f1 float(np.mean(f1_per_class)) if len(f1_per_class) else 0.0 LOGGER.info(fMean F1 Score: {mean_f1:.4f}) per_class_str [f{names[i]}: {f1_per_class[j]:.3f} for j, i in enumerate(class_indices)] LOGGER.info(fPer-class F1: {per_class_str}) return metrics, fitness注意BaseTrainer.validate()ultralytics/engine/trainer.py 第 865 行内部已用默认指标更新best_fitness因此在覆写时应先保存其旧值再做替换具体写法参见 docs/en/guides/custom-trainer.md 中“Saving the Best Model by Custom Metric”一节。层冻结freeze 参数的底层机制_setup_train()第 335-363 行对freeze参数的处理逻辑值得留意freeze可传整数冻结前 N 个层或层索引列表freeze_layer_names会追加.dfl等永久冻结层名知识蒸馏场景下还会追加teacher_model.若冻结后模型不再有任何可训练参数训练会直接抛出RuntimeError提示。因此做“先冻结主干、若干轮后再解冻”的迁移学习时既可以用model.train(freeze10, ...)这类内置参数也可以在子类__init__中通过self.add_callback(on_train_epoch_start, unfreeze_backbone)注册一个按 epoch 解冻的回调。完整的冻结/解冻实现示例见 docs/en/guides/custom-trainer.md。多卡训练的 SyncBatchNorm 与梯度裁剪在DetectionTrainer/RTDETRTrainer基础上还有两类常见覆写SyncBatchNorm小 batch 多卡训练时可在set_model_attributes()阶段模型已上 GPU、尚未被 DDP 包装的窗口期用nn.SyncBatchNorm.convert_sync_batchnorm(self.model)转换 BN梯度裁剪默认优化器步进在 ultralytics/engine/trainer.pyoptimizer_step()第 848 行处以max_norm10.0裁剪梯度DETR 系列模型通常需要更紧的0.1量级直接覆写该方法即可。RT-DETR 版本的完整代码与适用场景建议均收录于 docs/en/guides/custom-trainer.md。使用 YOLO 搭配自定义 TrainerYOLO模型类为 Trainer 家族提供了高层包装。利用这套架构可以在保持接口简洁的同时把训练过程的定制点下沉到自定义 Trainer 中from ultralytics import YOLO from ultralytics.models.yolo.detect import DetectionTrainer # Create a custom trainer class MyCustomTrainer(DetectionTrainer): def get_model(self, cfgNone, weightsNone, verboseTrue): Custom code implementation. # Initialize YOLO model model YOLO(yolo26n.pt) # Train with custom trainer results model.train(trainerMyCustomTrainer, datacoco8.yaml, epochs3)从源码看这一机制由 ultralytics/engine/model.py 中的train()方法第 777 行起支撑它接受trainer参数缺省时通过_smart_load(trainer)依据任务类型自动选择对应的 Trainer 子类当显式传入自定义 Trainer或注册了非默认回调时则会改用训练包装类并复用已加载的模型与权重第 851-861 行避免重复解析远程权重。其好处是显而易见的model.train(...)的调用方式与默认训练完全一致超参数、数据集路径、验证回调照常生效变化的只有“训练器内部如何执行”。适合把自定义 Trainer 封装成语义清晰的类供不同实验复用。其他引擎组件Validator 与 PredictorTrainer 并非唯一的引擎执行器。与BaseTrainer类似Validator验证器与Predictor预测器同样可以继承定制Validator对应 docs/en/reference/engine/validator.md负责验证集评估与指标计算可覆写以接入自定义评估逻辑Predictor对应 docs/en/reference/engine/predictor.md负责推理管线可覆写以定制预处理、后处理与结果输出。三者共享“继承 覆写钩子 回调”的统一设计哲学掌握 Trainer 的定制范式后可以平滑迁移到其他组件。推荐实践路径要深入掌握 Trainer 定制建议按以下顺序阅读仓库资料先阅读本主题的实战指南 docs/en/guides/custom-trainer.md其中涵盖自定义指标、类加权损失、自定义模型保存、主干冻结、分层学习率、SyncBatchNorm 与梯度裁剪等七类常见定制再对照阅读 BaseTrainer 源码参考 与 DetectionTrainer 参考确认每个被覆写方法的签名与调用时机需要扩展训练生命周期行为时查阅 Callbacks 指南优先用回调而非子类化实现轻量定制。FAQ 速查如何为特定任务定制 DetectionTrainer继承DetectionTrainer并重定义方法如get_model即可from ultralytics.models.yolo.detect import DetectionTrainer class CustomTrainer(DetectionTrainer): def get_model(self, cfgNone, weightsNone, verboseTrue): Loads a custom detection model given configuration and weight files. trainer CustomTrainer(overrides{...}) trainer.train() trained_model trainer.best # Get the best model如需修改损失函数或添加回调参见 Callbacks 指南。BaseTrainer 的关键组成有哪些除文档列出的get_model、get_dataloader、preprocess_batch、set_model_attributes、get_validator外还可覆写validate()、save_model()、build_optimizer()、label_loss_items()、optimizer_step()等方法源码见 ultralytics/engine/trainer.py 与 BaseTrainer 参考。如何给 DetectionTrainer 添加回调定义接收trainer的回调函数并用trainer.add_callback(event, func)注册。以每个 epoch 结束后记录权重路径为例from ultralytics.models.yolo.detect import DetectionTrainer # Callback to upload model weights def log_model(trainer): Logs the path of the last model weight used by the trainer. last_weight_path trainer.last print(last_weight_path) trainer DetectionTrainer(overrides{...}) trainer.add_callback(on_train_epoch_end, log_model) # Adds to existing callbacks trainer.train()DetectionTrainer支持非标准模型吗支持。DetectionTrainer高度灵活通过继承并覆写方法即可适配自定义模型典型即get_model。注意model.train(trainerMyCustomTrainer)传入的是自定义 Trainer 的类对象而非实例实例化由YOLO内部完成。【免费下载链接】ultralyticsUltralytics YOLO26, YOLO11, YOLOv8 — object detection, instance segmentation, semantic segmentation, image classification, pose estimation, object tracking项目地址: https://gitcode.com/GitHub_Trending/ul/ultralytics创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价