train_utility/tests/test_ultralytics_trainers.py
2026-08-04 11:14:51 +04:00

39 lines
1.4 KiB
Python

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