72 lines
2.3 KiB
Python
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]
|