423 lines
15 KiB
Python
423 lines
15 KiB
Python
from __future__ import annotations
|
|
|
|
import csv
|
|
import math
|
|
from pathlib import Path
|
|
from types import SimpleNamespace
|
|
from typing import Any
|
|
|
|
import pytest
|
|
|
|
from yolo_webui.config import YoloTask
|
|
from yolo_webui.mlflow_metrics import (
|
|
PER_CLASS_FIELDS,
|
|
TaskMetricsMonitor,
|
|
collect_monitoring_metrics,
|
|
collect_per_class_metrics,
|
|
)
|
|
|
|
|
|
def _component(offset: float = 0.0) -> SimpleNamespace:
|
|
return SimpleNamespace(
|
|
mp=0.75 + offset,
|
|
mr=0.65 + offset,
|
|
p=[0.8 + offset, 0.7 + offset],
|
|
r=[0.6 + offset, 0.7 + offset],
|
|
f1=[0.685714 + offset, 0.7 + offset],
|
|
map50=0.72 + offset,
|
|
map75=0.58 + offset,
|
|
map=0.61 + offset,
|
|
ap50=[0.76 + offset, 0.68 + offset],
|
|
ap=[0.64 + offset, 0.58 + offset],
|
|
ap_class_index=[0, 1],
|
|
)
|
|
|
|
|
|
def _metric_set(task: YoloTask) -> SimpleNamespace:
|
|
common: dict[str, Any] = {
|
|
"fitness": 0.61,
|
|
"speed": {
|
|
"preprocess": 0.1,
|
|
"inference": 1.2,
|
|
"loss": 0.3,
|
|
"postprocess": 0.4,
|
|
},
|
|
"names": {0: "cat", 1: "dog"},
|
|
"nt_per_class": [7, 5],
|
|
}
|
|
if task == "classify":
|
|
return SimpleNamespace(
|
|
**common,
|
|
top1=0.8,
|
|
top5=0.95,
|
|
confusion_matrix=SimpleNamespace(
|
|
matrix=[
|
|
[7, 2],
|
|
[1, 4],
|
|
],
|
|
names=common["names"],
|
|
),
|
|
)
|
|
|
|
common["box"] = _component()
|
|
if task == "segment":
|
|
common["seg"] = _component(0.05)
|
|
elif task == "pose":
|
|
common["pose"] = _component(0.1)
|
|
return SimpleNamespace(**common)
|
|
|
|
|
|
def _trainer(task: YoloTask, save_dir: Path | None = None) -> SimpleNamespace:
|
|
def label_loss_items(
|
|
loss: object,
|
|
prefix: str = "train",
|
|
) -> dict[str, float] | list[str]:
|
|
keys = [f"{prefix}/first_loss", f"{prefix}/second_loss"]
|
|
return keys if loss is None else dict(zip(keys, (0.25, 0.5)))
|
|
|
|
return SimpleNamespace(
|
|
metrics={
|
|
"val/first_loss": 0.2,
|
|
"val/second_loss": 0.3,
|
|
"val/future_non_loss_metric": 99.0,
|
|
},
|
|
validator=SimpleNamespace(metrics=_metric_set(task), training=True),
|
|
fitness=0.6,
|
|
best_fitness=0.65,
|
|
tloss=[0.25, 0.5],
|
|
label_loss_items=label_loss_items,
|
|
lr={"lr/pg0": 0.01, "lr/pg1": 0.02},
|
|
epoch_time=12.5,
|
|
save_dir=save_dir,
|
|
)
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("task", "required"),
|
|
[
|
|
(
|
|
"detect",
|
|
{
|
|
"monitor/quality/box/f1_at_optimal_confidence",
|
|
"monitor/quality/box/map75",
|
|
"monitor/quality/box/worst_class_map50_95",
|
|
},
|
|
),
|
|
(
|
|
"segment",
|
|
{
|
|
"monitor/quality/box/f1_at_optimal_confidence",
|
|
"monitor/quality/mask/f1_at_optimal_confidence",
|
|
"monitor/quality/mask/map75",
|
|
},
|
|
),
|
|
(
|
|
"pose",
|
|
{
|
|
"monitor/quality/box/f1_at_optimal_confidence",
|
|
"monitor/quality/keypoints/f1_at_optimal_confidence",
|
|
"monitor/quality/keypoints/map75",
|
|
},
|
|
),
|
|
(
|
|
"obb",
|
|
{
|
|
"monitor/quality/oriented_box/f1_at_optimal_confidence",
|
|
"monitor/quality/oriented_box/map75",
|
|
"monitor/quality/oriented_box/worst_class_map50_95",
|
|
},
|
|
),
|
|
(
|
|
"classify",
|
|
{
|
|
"monitor/quality/classification/top1_error",
|
|
"monitor/quality/classification/macro_f1",
|
|
"monitor/quality/classification/balanced_accuracy",
|
|
"monitor/quality/classification/worst_class_recall",
|
|
},
|
|
),
|
|
],
|
|
)
|
|
def test_collect_monitoring_metrics_has_task_specific_contract(
|
|
task: YoloTask,
|
|
required: set[str],
|
|
) -> None:
|
|
metrics = collect_monitoring_metrics(task, _trainer(task))
|
|
|
|
assert required <= metrics.keys()
|
|
assert metrics["monitor/fitness/ultralytics_current"] == pytest.approx(0.61)
|
|
assert metrics["monitor/fitness/ultralytics_best"] == pytest.approx(0.65)
|
|
expected_scale = 2 if task in {"segment", "pose"} else 1
|
|
assert metrics["monitor/fitness/task_score_current"] == pytest.approx(
|
|
0.61 / expected_scale
|
|
)
|
|
assert metrics["monitor/fitness/task_score_best"] == pytest.approx(
|
|
0.65 / expected_scale
|
|
)
|
|
assert metrics["monitor/loss/train_total"] == pytest.approx(0.75)
|
|
assert metrics["monitor/loss/validation_total"] == pytest.approx(0.5)
|
|
assert metrics["monitor/loss/generalization_gap"] == pytest.approx(-0.25)
|
|
assert metrics["monitor/loss/validation_to_train_ratio"] == pytest.approx(2 / 3)
|
|
assert metrics["monitor/optimization/learning_rate_mean"] == pytest.approx(0.015)
|
|
assert metrics["monitor/performance/epoch_seconds"] == pytest.approx(12.5)
|
|
assert metrics["monitor/performance/validation_inference_ms_per_image"] == pytest.approx(1.2)
|
|
|
|
|
|
def test_classification_metrics_are_derived_from_confusion_matrix() -> None:
|
|
metrics = collect_monitoring_metrics("classify", _trainer("classify"))
|
|
|
|
assert metrics["monitor/quality/classification/top1_error"] == pytest.approx(0.2)
|
|
assert metrics["monitor/quality/classification/macro_precision"] == pytest.approx(
|
|
(7 / 9 + 4 / 5) / 2
|
|
)
|
|
assert metrics["monitor/quality/classification/macro_recall"] == pytest.approx(
|
|
(7 / 8 + 4 / 6) / 2
|
|
)
|
|
assert metrics["monitor/quality/classification/classes_evaluated"] == 2
|
|
|
|
|
|
def test_mask_and_keypoint_metrics_use_their_own_components() -> None:
|
|
segment_metrics = collect_monitoring_metrics("segment", _trainer("segment"))
|
|
pose_metrics = collect_monitoring_metrics("pose", _trainer("pose"))
|
|
|
|
assert segment_metrics[
|
|
"monitor/quality/mask/f1_at_optimal_confidence"
|
|
] == pytest.approx((0.735714 + 0.75) / 2)
|
|
assert pose_metrics[
|
|
"monitor/quality/keypoints/f1_at_optimal_confidence"
|
|
] == pytest.approx((0.785714 + 0.8) / 2)
|
|
assert segment_metrics[
|
|
"monitor/quality/mask/f1_at_optimal_confidence"
|
|
] != segment_metrics["monitor/quality/box/f1_at_optimal_confidence"]
|
|
|
|
|
|
def test_classification_macro_metrics_include_false_positive_only_class() -> None:
|
|
trainer = _trainer("classify")
|
|
trainer.validator.metrics.confusion_matrix = SimpleNamespace(
|
|
matrix=[
|
|
[7, 2, 0],
|
|
[1, 4, 0],
|
|
[1, 0, 0],
|
|
],
|
|
names={0: "cat", 1: "dog", 2: "fox"},
|
|
)
|
|
|
|
metrics = collect_monitoring_metrics("classify", trainer)
|
|
rows = collect_per_class_metrics("classify", trainer)
|
|
|
|
assert metrics["monitor/quality/classification/classes_evaluated"] == 3
|
|
assert metrics["monitor/quality/classification/macro_precision"] == pytest.approx(
|
|
(7 / 9 + 4 / 5 + 0.0) / 3
|
|
)
|
|
assert metrics["monitor/quality/classification/macro_recall"] == pytest.approx(
|
|
(7 / 9 + 4 / 6) / 2
|
|
)
|
|
assert rows[-1]["class_name"] == "fox"
|
|
assert rows[-1]["support"] == 0
|
|
assert rows[-1]["precision"] == 0
|
|
|
|
|
|
def test_monitor_drops_non_finite_and_non_numeric_values() -> None:
|
|
trainer = _trainer("detect")
|
|
trainer.best_fitness = float("nan")
|
|
trainer.epoch_time = float("inf")
|
|
trainer.lr = {"lr/pg0": "not-a-number"}
|
|
trainer.validator.metrics.box.map75 = None
|
|
|
|
metrics = collect_monitoring_metrics("detect", trainer)
|
|
|
|
assert "monitor/fitness/ultralytics_best" not in metrics
|
|
assert "monitor/fitness/task_score_best" not in metrics
|
|
assert "monitor/performance/epoch_seconds" not in metrics
|
|
assert "monitor/optimization/learning_rate_mean" not in metrics
|
|
assert "monitor/quality/box/map75" not in metrics
|
|
assert all(
|
|
value == value and value not in (float("inf"), float("-inf"))
|
|
for value in metrics.values()
|
|
)
|
|
|
|
|
|
def test_callback_enriches_trainer_metrics_before_mlflow_logging() -> None:
|
|
trainer = _trainer("obb")
|
|
monitor = TaskMetricsMonitor("obb", mlflow_enabled=True)
|
|
|
|
monitor.on_fit_epoch_end(trainer)
|
|
|
|
assert trainer.metrics["monitor/fitness/ultralytics_current"] == pytest.approx(0.61)
|
|
assert trainer.metrics[
|
|
"monitor/quality/oriented_box/f1_at_optimal_confidence"
|
|
] == pytest.approx((0.685714 + 0.7) / 2)
|
|
|
|
|
|
def test_final_validation_does_not_repeat_stale_train_diagnostics() -> None:
|
|
trainer = _trainer("detect")
|
|
trainer.validator.training = False
|
|
|
|
metrics = collect_monitoring_metrics("detect", trainer)
|
|
|
|
assert "monitor/quality/box/f1_at_optimal_confidence" in metrics
|
|
assert "monitor/performance/validation_inference_ms_per_image" in metrics
|
|
assert "monitor/loss/train_total" not in metrics
|
|
assert "monitor/loss/validation_total" not in metrics
|
|
assert "monitor/optimization/learning_rate_mean" not in metrics
|
|
assert "monitor/performance/epoch_seconds" not in metrics
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("task", "expected_components"),
|
|
[
|
|
("detect", {"box"}),
|
|
("segment", {"box", "mask"}),
|
|
("pose", {"box", "keypoints"}),
|
|
("obb", {"oriented_box"}),
|
|
("classify", {"classification"}),
|
|
],
|
|
)
|
|
def test_final_per_class_rows_cover_every_task_component(
|
|
task: YoloTask,
|
|
expected_components: set[str],
|
|
) -> None:
|
|
rows = collect_per_class_metrics(task, _trainer(task))
|
|
|
|
assert {row["component"] for row in rows} == expected_components
|
|
assert {row["class_name"] for row in rows} == {"cat", "dog"}
|
|
assert all(row["support"] > 0 for row in rows)
|
|
|
|
|
|
def test_per_class_rows_preserve_positions_of_non_finite_values() -> None:
|
|
trainer = _trainer("detect")
|
|
component = trainer.validator.metrics.box
|
|
component.ap_class_index = [0, 1, 2]
|
|
component.p = [0.8, float("nan"), 0.6]
|
|
component.r = [0.7, 0.5, 0.4]
|
|
component.f1 = [0.746, 0.5, 0.48]
|
|
component.ap50 = [0.75, 0.55, 0.45]
|
|
component.ap = [0.65, 0.45, 0.35]
|
|
trainer.validator.metrics.names = {0: "first", 1: "second", 2: "third"}
|
|
trainer.validator.metrics.nt_per_class = [5, 4, 3]
|
|
|
|
rows = collect_per_class_metrics("detect", trainer)
|
|
|
|
assert [row["class_id"] for row in rows] == [0, 1, 2]
|
|
assert rows[1]["precision"] == ""
|
|
assert rows[2]["precision"] == pytest.approx(0.6)
|
|
|
|
|
|
def test_train_end_writes_mlflow_uploadable_per_class_csv(tmp_path: Path) -> None:
|
|
trainer = _trainer("segment", tmp_path)
|
|
monitor = TaskMetricsMonitor("segment", mlflow_enabled=True)
|
|
|
|
monitor.on_train_end(trainer)
|
|
|
|
output = tmp_path / "monitoring" / "task_metrics.csv"
|
|
assert output.exists()
|
|
with output.open(encoding="utf-8") as stream:
|
|
reader = csv.DictReader(stream)
|
|
rows = list(reader)
|
|
assert tuple(reader.fieldnames or ()) == PER_CLASS_FIELDS
|
|
assert len(rows) == 4
|
|
assert {row["component"] for row in rows} == {"box", "mask"}
|
|
assert {row["task"] for row in rows} == {"segment"}
|
|
assert {row["class_id"] for row in rows} == {"0", "1"}
|
|
assert {row["class_name"] for row in rows} == {"cat", "dog"}
|
|
for row in rows:
|
|
for field in ("support", "precision", "recall", "f1", "map50", "map50_95"):
|
|
assert math.isfinite(float(row[field]))
|
|
|
|
|
|
def test_disabled_mlflow_does_not_add_metrics_or_write_artifact(tmp_path: Path) -> None:
|
|
trainer = _trainer("detect", tmp_path)
|
|
original_metrics = dict(trainer.metrics)
|
|
monitor = TaskMetricsMonitor("detect", mlflow_enabled=False)
|
|
|
|
monitor.on_fit_epoch_end(trainer)
|
|
monitor.on_train_end(trainer)
|
|
|
|
assert trainer.metrics == original_metrics
|
|
assert not (tmp_path / "monitoring" / "task_metrics.csv").exists()
|
|
|
|
|
|
def test_train_end_writes_header_when_no_classes_were_evaluated(tmp_path: Path) -> None:
|
|
trainer = _trainer("detect", tmp_path)
|
|
component = trainer.validator.metrics.box
|
|
for attribute in ("p", "r", "f1", "ap50", "ap", "ap_class_index"):
|
|
setattr(component, attribute, [])
|
|
|
|
TaskMetricsMonitor("detect", mlflow_enabled=True).on_train_end(trainer)
|
|
|
|
output = tmp_path / "monitoring" / "task_metrics.csv"
|
|
with output.open(encoding="utf-8") as stream:
|
|
reader = csv.DictReader(stream)
|
|
rows = list(reader)
|
|
assert tuple(reader.fieldnames or ()) == PER_CLASS_FIELDS
|
|
assert rows == []
|
|
|
|
|
|
def test_monitoring_contract_is_persisted_by_mlflow(
|
|
tmp_path: Path,
|
|
monkeypatch: pytest.MonkeyPatch,
|
|
) -> None:
|
|
import mlflow
|
|
from mlflow.tracking import MlflowClient
|
|
from ultralytics.utils.callbacks import mlflow as ultralytics_mlflow
|
|
from yolo_webui.ultralytics_trainers import trainer_for_task
|
|
|
|
previous_uri = mlflow.get_tracking_uri()
|
|
tracking_uri = f"sqlite:///{tmp_path / 'mlflow.db'}"
|
|
run_id = ""
|
|
try:
|
|
monkeypatch.setenv("YOLO_WEBUI_MLFLOW_RUN_GROUP", "test-group")
|
|
mlflow.set_tracking_uri(tracking_uri)
|
|
client = MlflowClient(tracking_uri=tracking_uri)
|
|
experiment_id = client.create_experiment(
|
|
"task-monitor-test",
|
|
artifact_location=(tmp_path / "artifacts").as_uri(),
|
|
)
|
|
trainer_class = trainer_for_task("obb")
|
|
trainer = trainer_class.__new__(trainer_class)
|
|
trainer.__dict__.update(vars(_trainer("obb", tmp_path / "run")))
|
|
trainer._mlflow_active = True
|
|
trainer.epoch = 0
|
|
trainer.callbacks = {
|
|
"on_train_start": [],
|
|
"on_fit_epoch_end": [ultralytics_mlflow.on_fit_epoch_end],
|
|
"on_train_end": [],
|
|
}
|
|
trainer._install_task_metrics_monitor()
|
|
trainer._task_metrics_monitor.mlflow_enabled = True
|
|
|
|
with mlflow.start_run(experiment_id=experiment_id, run_name="obb-contract") as run:
|
|
run_id = run.info.run_id
|
|
for callback in trainer.callbacks["on_train_start"]:
|
|
callback(trainer)
|
|
for callback in trainer.callbacks["on_fit_epoch_end"]:
|
|
callback(trainer)
|
|
|
|
trainer.epoch = 1
|
|
trainer.validator.training = False
|
|
trainer.metrics = {"metrics/mAP50": 0.72}
|
|
for callback in trainer.callbacks["on_fit_epoch_end"]:
|
|
callback(trainer)
|
|
for callback in trainer.callbacks["on_train_end"]:
|
|
callback(trainer)
|
|
|
|
persisted = client.get_run(run_id)
|
|
history = client.get_metric_history(
|
|
run_id,
|
|
"monitor/quality/oriented_box/f1_at_optimal_confidence",
|
|
)
|
|
assert persisted.data.tags["yolo.task"] == "obb"
|
|
assert persisted.data.tags["monitoring.schema_version"] == "1"
|
|
assert persisted.data.tags["yolo.run_group"] == "test-group"
|
|
assert persisted.data.metrics[
|
|
"monitor/fitness/ultralytics_current"
|
|
] == pytest.approx(0.61)
|
|
assert [point.step for point in history] == [0, 1]
|
|
assert {
|
|
artifact.path
|
|
for artifact in client.list_artifacts(run_id, "monitoring")
|
|
} == {"monitoring/task_metrics.csv"}
|
|
finally:
|
|
if mlflow.active_run() is not None:
|
|
mlflow.end_run(status="KILLED")
|
|
mlflow.set_tracking_uri(previous_uri)
|