高级自定义#
Ultralytics YOLO 的命令行和 Python 接口都是基于底层引擎执行器构建的高级抽象。本指南将重点介绍 Trainer 引擎,并说明如何根据具体需求对其进行自定义。
观看: 精通 Ultralytics YOLO:高级自定义
有关常见训练器自定义的实用示例——自定义指标、类别加权损失、模型保存、冻结骨干网络以及逐层学习率——请参阅自定义训练器指南。
BaseTrainer#
BaseTrainer 类提供通用训练流程,可适配各种任务。你可以在遵循所需格式的前提下,重写特定函数或操作来自定义它。例如,重写以下函数即可集成自己的自定义模型和数据加载器:
get_model(cfg, weights):构建待训练的模型。get_dataloader(dataset_path, batch_size, rank, mode):构建数据加载器。
如需了解更多详情和源代码,请参阅 BaseTrainer 参考文档。
DetectionTrainer#
以下介绍如何使用和自定义 Ultralytics YOLO DetectionTrainer:
from ultralytics.models.yolo.detect import DetectionTrainer
trainer = DetectionTrainer(overrides={...})
trainer.train()
trained_model = trainer.best # 获取最佳模型自定义 DetectionTrainer#
要训练尚未直接支持的自定义检测模型,请重载现有的 get_model 功能:
from ultralytics.models.yolo.detect import DetectionTrainer
class CustomTrainer(DetectionTrainer):
def get_model(self, cfg=None, weights=None, verbose=True):
"""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 a custom loss function for the model."""
class CustomTrainer(DetectionTrainer):
def get_model(self, cfg=None, weights=None, verbose=True):
"""Returns a customized detection model instance configured with specified config and weights."""
return MyCustomModel(...)
# 用于记录模型权重的回调
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) # 添加到现有回调中
trainer.train()如需了解更多有关回调触发事件和入口点的信息,请参阅回调指南。
其他引擎组件#
你可以用类似的方法自定义其他组件,例如 Validators 和 Predictors。如需了解更多信息,请参阅验证器和预测器文档。
将 YOLO 与自定义训练器配合使用#
YOLO 模型类为训练器类提供了高级封装。你可以利用这一架构,让机器学习工作流更加灵活:
from ultralytics import YOLO
from ultralytics.models.yolo.detect import DetectionTrainer
# 创建自定义训练器
class MyCustomTrainer(DetectionTrainer):
def get_model(self, cfg=None, weights=None, verbose=True):
"""Custom code implementation."""
# 初始化 YOLO 模型
model = YOLO("yolo26n.pt")
# 使用自定义训练器进行训练
results = model.train(trainer=MyCustomTrainer, data="coco8.yaml", epochs=3)这种方法让你既能保留 YOLO 接口的简洁性,又能自定义底层训练过程,以满足具体需求。
常见问题#
要针对特定任务自定义
DetectionTrainer,可以重写它的方法,使其适配你的自定义模型和数据加载器。首先继承DetectionTrainer,然后重新定义get_model等方法,以实现自定义功能。示例如下:from ultralytics.models.yolo.detect import DetectionTrainer class CustomTrainer(DetectionTrainer): def get_model(self, cfg=None, weights=None, verbose=True): """Loads a custom detection model given configuration and weight files.""" trainer = CustomTrainer(overrides={...}) trainer.train() trained_model = trainer.best # 获取最佳模型BaseTrainer是训练流程的基础,可通过重写其通用方法来适配各种任务。关键组件包括:get_model(cfg, weights):构建待训练的模型。get_dataloader(dataset_path, batch_size, rank, mode):构建数据加载器。preprocess_batch():在模型前向传播之前处理批次预处理。set_model_attributes():根据数据集信息设置模型属性。get_validator():返回用于模型评估的验证器。
如需了解更多有关自定义和源代码的信息,请参阅
BaseTrainer参考文档。在
DetectionTrainer中添加回调,以监控并修改训练过程。以下介绍如何添加回调,在每个训练轮次后记录模型权重:from ultralytics.models.yolo.detect import DetectionTrainer # 用于记录模型权重的回调 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) # 添加到现有回调中 trainer.train()如需详细了解回调事件和入口点,请参阅回调指南。
Ultralytics YOLO 为强大的引擎执行器提供高级抽象,非常适合快速开发和自定义。主要优势包括:
- 易于使用:命令行和 Python 接口都能简化复杂任务。
- 性能:针对实时目标检测和各种视觉 AI 应用进行了优化。
- 可自定义:可轻松扩展以支持自定义模型、损失函数和数据加载器。
- 模块化:可以独立修改各个组件,而不会影响整个流程。
- 集成:可与 ML 生态系统中的主流框架和工具无缝协作。
探索主要的 Ultralytics YOLO 页面,进一步了解 YOLO 的功能。
可以,
DetectionTrainer非常灵活,可针对非标准模型进行自定义。继承DetectionTrainer并重载方法,以满足特定模型的需求。以下是一个简单示例:from ultralytics.models.yolo.detect import DetectionTrainer class CustomDetectionTrainer(DetectionTrainer): def get_model(self, cfg=None, weights=None, verbose=True): """Loads a custom detection model.""" trainer = CustomDetectionTrainer(overrides={...}) trainer.train()如需完整的操作说明和示例,请参阅
DetectionTrainer参考文档。