YOLO Vision 2026:

高级自定义#

Ultralytics YOLO 的命令行和 Python 接口都是构建在基础引擎执行器之上的高级抽象。本指南专注于 Trainer 引擎,解释如何针对你的具体需求对其进行自定义。



Watch: Mastering Ultralytics YOLO: Advanced Customization
提示

有关常见训练器自定义的实际示例——自定义指标、类别加权损失、模型保存、主干网络冻结和逐层学习率——请参见自定义训练器指南。

BaseTrainer#

BaseTrainer 类提供了一个通用的训练例程,适用于各种任务。通过重写特定函数或操作并在遵守所需格式的同时进行自定义。例如,通过重写以下函数集成你自己的自定义模型和数据加载器:

  • get_model(cfg, weights):构建要训练的模型。
  • get_dataloader():构建数据加载器。

有关更多详细信息和源代码,请参见BaseTrainer 参考

DetectionTrainer#

以下是如何使用和自定义 Ultralytics YOLO DetectionTrainer

from ultralytics.models.yolo.detect import DetectionTrainer

trainer = DetectionTrainer(overrides={...})
trainer.train()
trained_model = trainer.best  # Get the best model

自定义 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()

通过修改损失函数或添加回调来进一步自定义训练器,以便每 10 个轮次将模型上传到 Google Drive。这是一个例子:

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, cfg=None, weights=None, verbose=True):
        """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()

有关回调触发事件和切入点的更多信息,请参见回调指南

其他引擎组件#

以类似方式自定义其他组件,如 ValidatorsPredictors。有关更多信息,请参考验证器预测器的文档。

将 YOLO 与自定义训练器结合使用#

YOLO 模型类为 Trainer 类提供了高级包装器。你可以利用此架构在你的机器学习工作流中获得更大的灵活性:

from ultralytics import YOLO
from ultralytics.models.yolo.detect import DetectionTrainer

# Create a custom trainer
class MyCustomTrainer(DetectionTrainer):
    def get_model(self, cfg=None, weights=None, verbose=True):
        """Custom code implementation."""

# Initialize YOLO model
model = YOLO("yolo26n.pt")

# Train with custom trainer
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  # Get the best model

    有关进一步自定义(例如更改损失函数或添加回调),请参考回调指南

  • BaseTrainer 作为训练例程的基础,可以通过重写其通用方法来针对各种任务进行自定义。核心组件包括:

    • get_model(cfg, weights):构建要训练的模型。
    • get_dataloader():构建数据加载器。
    • preprocess_batch():在模型前向传播之前处理批次预处理。
    • set_model_attributes():根据数据集信息设置模型属性。
    • get_validator():返回用于模型评估的验证器。

    有关自定义和源代码的更多详细信息,请参见BaseTrainer 参考

  • DetectionTrainer 中添加回调以监控和修改训练过程。以下是如何添加一个回调以在每个训练轮次后记录模型权重的示例:

    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()

    有关回调事件和切入点的更多详细信息,请参考回调指南

  • 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 参考

评论