from __future__ import annotations from typing import Any, ClassVar from ultralytics.models.yolo.classify.train import ClassificationTrainer from ultralytics.models.yolo.detect.train import DetectionTrainer from ultralytics.models.yolo.obb.train import OBBTrainer from ultralytics.models.yolo.pose.train import PoseTrainer from ultralytics.models.yolo.segment.train import SegmentationTrainer from ultralytics.utils import SETTINGS from .config import YoloTask from .mlflow_metrics import TaskMetricsMonitor class _TaskMetricsTrainerMixin: """Install monitoring inside the trainer so Ultralytics DDP keeps it.""" monitoring_task: ClassVar[YoloTask] def __init__(self, *args: Any, **kwargs: Any) -> None: super().__init__(*args, **kwargs) self._install_task_metrics_monitor() def _install_task_metrics_monitor(self) -> None: monitor = TaskMetricsMonitor( self.monitoring_task, mlflow_enabled=SETTINGS["mlflow"] is True, ) self._task_metrics_monitor = monitor callbacks = ( ("on_train_start", monitor.on_train_start), ("on_fit_epoch_end", monitor.on_fit_epoch_end), ("on_train_end", monitor.on_train_end), ) for event, callback in callbacks: # BaseTrainer adds integrations in __init__. Prepending ensures our # metrics and CSV exist before Ultralytics logs them to MLflow. self.callbacks.setdefault(event, []).insert(0, callback) class MonitoredDetectionTrainer(_TaskMetricsTrainerMixin, DetectionTrainer): monitoring_task = "detect" class MonitoredSegmentationTrainer(_TaskMetricsTrainerMixin, SegmentationTrainer): monitoring_task = "segment" class MonitoredClassificationTrainer(_TaskMetricsTrainerMixin, ClassificationTrainer): monitoring_task = "classify" class MonitoredPoseTrainer(_TaskMetricsTrainerMixin, PoseTrainer): monitoring_task = "pose" class MonitoredOBBTrainer(_TaskMetricsTrainerMixin, OBBTrainer): monitoring_task = "obb" MONITORED_TRAINERS = { "detect": MonitoredDetectionTrainer, "segment": MonitoredSegmentationTrainer, "classify": MonitoredClassificationTrainer, "pose": MonitoredPoseTrainer, "obb": MonitoredOBBTrainer, } def trainer_for_task(task: YoloTask) -> type: return MONITORED_TRAINERS[task]