39 lines
1.4 KiB
Python
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
|