from __future__ import annotations from typing import Any import pytest from yolo_webui.config import YoloTask from yolo_webui.mlflow_metrics import TaskMetricsMonitor from yolo_webui.ultralytics_trainers import MONITORED_TRAINERS, trainer_for_task @pytest.mark.parametrize("task", ["detect", "segment", "classify", "pose", "obb"]) def test_every_task_uses_an_importable_monitored_trainer(task: YoloTask) -> None: trainer_class = trainer_for_task(task) assert trainer_class is MONITORED_TRAINERS[task] assert trainer_class.monitoring_task == task assert trainer_class.__module__ == "yolo_webui.ultralytics_trainers" def test_monitor_callbacks_are_prepended_before_integrations() -> None: trainer_class = trainer_for_task("detect") trainer = trainer_class.__new__(trainer_class) integration_callbacks: dict[str, Any] = { "on_train_start": lambda _trainer: None, "on_fit_epoch_end": lambda _trainer: None, "on_train_end": lambda _trainer: None, } trainer.callbacks = { event: [callback] for event, callback in integration_callbacks.items() } trainer._install_task_metrics_monitor() assert isinstance(trainer._task_metrics_monitor, TaskMetricsMonitor) for event, integration_callback in integration_callbacks.items(): assert trainer.callbacks[event][0].__self__ is trainer._task_metrics_monitor assert trainer.callbacks[event][1] is integration_callback