208 lines
6.7 KiB
Python
Executable file
208 lines
6.7 KiB
Python
Executable file
from __future__ import annotations
|
|
|
|
import sys
|
|
import signal
|
|
from pathlib import Path
|
|
from types import ModuleType, SimpleNamespace
|
|
from typing import Any
|
|
|
|
import pytest
|
|
|
|
from yolo_webui.config import MlflowConfig, TrainingConfig
|
|
from yolo_webui.trainer import TrainingEvent, TrainingRunner
|
|
|
|
|
|
class FakeProcess:
|
|
def __init__(self) -> None:
|
|
self.signals: list[int] = []
|
|
self.terminate_calls = 0
|
|
self.kill_calls = 0
|
|
|
|
def send_signal(self, signum: int) -> None:
|
|
self.signals.append(signum)
|
|
|
|
def terminate(self) -> None:
|
|
self.terminate_calls += 1
|
|
|
|
def kill(self) -> None:
|
|
self.kill_calls += 1
|
|
|
|
def poll(self) -> None:
|
|
return None
|
|
|
|
|
|
def test_runner_wires_yolo_callbacks_and_returns_output(
|
|
monkeypatch: Any, tmp_path: Path
|
|
) -> None:
|
|
settings_updates: list[dict[str, bool]] = []
|
|
constructed: list[tuple[str, str]] = []
|
|
train_arguments: list[dict[str, Any]] = []
|
|
|
|
class FakeSettings:
|
|
def __init__(self) -> None:
|
|
self.values = {"mlflow": True}
|
|
|
|
def __getitem__(self, key: str) -> bool:
|
|
return self.values[key]
|
|
|
|
def update(self, values: dict[str, bool]) -> None:
|
|
settings_updates.append(values)
|
|
self.values.update(values)
|
|
|
|
class FakeYOLO:
|
|
def __init__(self, model: str, task: str) -> None:
|
|
constructed.append((model, task))
|
|
self.callbacks: dict[str, list[Any]] = {}
|
|
self.trainer = SimpleNamespace(
|
|
args=SimpleNamespace(epochs=2),
|
|
epoch=0,
|
|
metrics={},
|
|
fitness=0.0,
|
|
best_fitness=0.0,
|
|
stop=False,
|
|
save_dir=tmp_path / "run",
|
|
)
|
|
|
|
def add_callback(self, name: str, callback: Any) -> None:
|
|
self.callbacks.setdefault(name, []).append(callback)
|
|
|
|
def run_callbacks(self, name: str) -> None:
|
|
for callback in self.callbacks.get(name, []):
|
|
callback(self.trainer)
|
|
|
|
def train(self, **kwargs: Any) -> None:
|
|
train_arguments.append(kwargs)
|
|
self.run_callbacks("on_train_start")
|
|
for epoch in range(2):
|
|
self.trainer.epoch = epoch
|
|
self.trainer.metrics = {"metrics/mAP50": 0.5 + epoch / 10}
|
|
self.trainer.fitness = 0.5 + epoch / 10
|
|
self.trainer.best_fitness = self.trainer.fitness
|
|
self.run_callbacks("on_fit_epoch_end")
|
|
self.run_callbacks("on_train_end")
|
|
|
|
fake_ultralytics = ModuleType("ultralytics")
|
|
fake_ultralytics.YOLO = FakeYOLO # type: ignore[attr-defined]
|
|
fake_ultralytics.settings = FakeSettings() # type: ignore[attr-defined]
|
|
fake_trainers = ModuleType("yolo_webui.ultralytics_trainers")
|
|
|
|
class FakeTrainer:
|
|
pass
|
|
|
|
fake_trainers.trainer_for_task = lambda _task: FakeTrainer # type: ignore[attr-defined]
|
|
monkeypatch.setitem(sys.modules, "ultralytics", fake_ultralytics)
|
|
monkeypatch.setitem(sys.modules, "yolo_webui.ultralytics_trainers", fake_trainers)
|
|
|
|
config = TrainingConfig(
|
|
dataset="dataset.yaml",
|
|
model="model.pt",
|
|
task="pose",
|
|
epochs=2,
|
|
mlflow=MlflowConfig(enabled=False),
|
|
)
|
|
events: list[TrainingEvent] = []
|
|
|
|
output = TrainingRunner().train(config, events.append)
|
|
|
|
assert output == tmp_path / "run"
|
|
assert constructed == [("models/model.pt", "pose")]
|
|
assert [u for u in settings_updates if "mlflow" in u] == [{"mlflow": False}, {"mlflow": True}]
|
|
assert train_arguments[0]["data"] == "dataset.yaml"
|
|
assert train_arguments[0]["verbose"] is True
|
|
assert train_arguments[0]["trainer"] is FakeTrainer
|
|
assert [event.kind for event in events] == ["info", "started", "epoch", "epoch", "success"]
|
|
assert "mAP50=0.5" in events[2].message
|
|
assert "mAP50=0.6" in events[3].message
|
|
|
|
|
|
def test_ultralytics_mlflow_setting_is_restored_when_model_loading_fails(
|
|
monkeypatch: Any,
|
|
) -> None:
|
|
settings_updates: list[bool] = []
|
|
|
|
class FakeSettings:
|
|
def __init__(self) -> None:
|
|
self.mlflow = True
|
|
|
|
def __getitem__(self, key: str) -> bool:
|
|
assert key == "mlflow"
|
|
return self.mlflow
|
|
|
|
def update(self, values: dict[str, bool]) -> None:
|
|
self.mlflow = values["mlflow"]
|
|
settings_updates.append(self.mlflow)
|
|
|
|
class FailingYOLO:
|
|
def __init__(self, _model: str, task: str) -> None:
|
|
assert task == "detect"
|
|
raise RuntimeError("model load failed")
|
|
|
|
fake_ultralytics = ModuleType("ultralytics")
|
|
fake_ultralytics.YOLO = FailingYOLO # type: ignore[attr-defined]
|
|
fake_ultralytics.settings = FakeSettings() # type: ignore[attr-defined]
|
|
fake_trainers = ModuleType("yolo_webui.ultralytics_trainers")
|
|
fake_trainers.trainer_for_task = lambda _task: object # type: ignore[attr-defined]
|
|
monkeypatch.setitem(sys.modules, "ultralytics", fake_ultralytics)
|
|
monkeypatch.setitem(sys.modules, "yolo_webui.ultralytics_trainers", fake_trainers)
|
|
config = TrainingConfig(
|
|
dataset="dataset.yaml",
|
|
model="model.pt",
|
|
mlflow=MlflowConfig(enabled=False),
|
|
)
|
|
|
|
with pytest.raises(RuntimeError, match="model load failed"):
|
|
TrainingRunner().train(config, lambda _event: None)
|
|
|
|
assert settings_updates == [False, True]
|
|
assert fake_ultralytics.settings["mlflow"] is True # type: ignore[index]
|
|
|
|
|
|
def test_final_checkpoint_validation_is_not_reported_as_an_extra_epoch() -> None:
|
|
events: list[TrainingEvent] = []
|
|
trainer = SimpleNamespace(
|
|
args=SimpleNamespace(epochs=1),
|
|
epoch=1,
|
|
metrics={"metrics/mAP50": 0.75},
|
|
validator=SimpleNamespace(training=False),
|
|
)
|
|
|
|
TrainingRunner()._on_epoch_end(events.append)(trainer)
|
|
|
|
assert events == []
|
|
|
|
|
|
def test_prepare_run_clears_previous_stop_request() -> None:
|
|
runner = TrainingRunner()
|
|
runner.request_stop()
|
|
assert runner.stop_requested is True
|
|
|
|
runner.prepare_run()
|
|
|
|
assert runner.stop_requested is False
|
|
|
|
|
|
def test_early_stop_is_delivered_after_subprocess_ready() -> None:
|
|
runner = TrainingRunner()
|
|
process = FakeProcess()
|
|
runner.request_stop()
|
|
|
|
runner.set_subprocess(process, ready=False)
|
|
assert process.signals == []
|
|
|
|
runner.mark_subprocess_ready()
|
|
|
|
assert process.signals == [signal.SIGTERM]
|
|
assert process.terminate_calls == 0
|
|
runner.clear_subprocess()
|
|
|
|
|
|
def test_ready_subprocess_receives_cooperative_signal_not_terminate() -> None:
|
|
runner = TrainingRunner()
|
|
process = FakeProcess()
|
|
runner.set_subprocess(process, ready=True)
|
|
|
|
runner.request_stop()
|
|
|
|
assert process.signals == [signal.SIGTERM]
|
|
assert process.terminate_calls == 0
|
|
runner.clear_subprocess()
|