train_utility/src/yolo_webui/ultralytics_trainers.py
2026-08-04 11:14:51 +04:00

72 lines
2.3 KiB
Python

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]