Expand application functionality and coverage
This commit is contained in:
parent
537ef2e489
commit
53f758cc07
11 changed files with 2213 additions and 58 deletions
32
README.md
32
README.md
|
|
@ -96,6 +96,31 @@ uv run mlflow ui --backend-store-uri sqlite:///mlflow.db
|
|||
артефакты запуска, а не версии MLflow Model Registry: raw YOLO checkpoint не имеет
|
||||
стандартной MLflow `MLmodel`-упаковки.
|
||||
|
||||
WebUI дополняет штатные метрики стабильным namespace `monitor/*`. Значения
|
||||
собираются после validation текущей эпохи, поэтому не запаздывают на одну эпоху:
|
||||
|
||||
| Задача | Основные дополнительные ряды |
|
||||
|---|---|
|
||||
| `detect` | box F1 при оптимальном confidence, mAP@0.75, метрики худшего класса |
|
||||
| `segment` | F1 при оптимальном confidence и mAP@0.75 отдельно для box и mask |
|
||||
| `classify` | top-1/top-5 error, macro precision/recall/F1, weighted F1, balanced accuracy |
|
||||
| `pose` | F1 при оптимальном confidence и mAP@0.75 отдельно для box и keypoints (OKS) |
|
||||
| `obb` | F1 при оптимальном confidence, mAP@0.75 и метрики худшего класса для oriented boxes |
|
||||
|
||||
Для всех задач также записываются исходный Ultralytics fitness и нормализованный
|
||||
`task_score` (для `segment`/`pose` сумма box+mask/keypoints приводится к диапазону
|
||||
0–1), суммарные train/validation loss, gap/ratio между ними, средний learning rate,
|
||||
время эпохи и скорость validation по стадиям. Теги `yolo.task` и
|
||||
`monitoring.schema_version` позволяют фильтровать совместимые запуски. Финальные
|
||||
per-class precision/recall/F1/AP и support сохраняются в артефакте
|
||||
`monitoring/task_metrics.csv`, чтобы не создавать сотни поэпоховых рядов для
|
||||
датасетов с большим числом классов. Monitor подключается внутри task-specific
|
||||
trainer и сохраняется при запуске Ultralytics DDP на нескольких GPU.
|
||||
|
||||
Шаги MLflow остаются совместимыми с Ultralytics: первая эпоха имеет `step=0`, а
|
||||
финальная проверка лучшего checkpoint добавляется отдельной точкой после последней
|
||||
эпохи. В WebUI эпохи отображаются привычно, начиная с 1.
|
||||
|
||||
Проверка интеграции на минимальных датасетах для всех пяти задач:
|
||||
|
||||
```bash
|
||||
|
|
@ -103,8 +128,13 @@ uv run scripts/run_yolo26_smoke_training.py --mlflow
|
|||
uv run scripts/verify_mlflow_smoke.py
|
||||
```
|
||||
|
||||
Smoke-runner присваивает всей пятёрке запусков уникальный тег `yolo.run_group`;
|
||||
проверка берёт только самый свежий batch и не смешивает его со старыми успешными
|
||||
run-ами.
|
||||
|
||||
Второй скрипт завершается с ошибкой, если отсутствует experiment/run, параметры,
|
||||
метрики, `results.csv`, `best.pt` или `last.pt` хотя бы для одной задачи.
|
||||
task-specific метрики, их история, теги, `monitoring/task_metrics.csv`,
|
||||
`results.csv`, `best.pt` или `last.pt` хотя бы для одной задачи.
|
||||
|
||||
## Проверка
|
||||
|
||||
|
|
|
|||
|
|
@ -4,8 +4,10 @@ from __future__ import annotations
|
|||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import traceback
|
||||
import uuid
|
||||
from pathlib import Path
|
||||
|
||||
from yolo_webui import TrainingConfig, TrainingRunner
|
||||
|
|
@ -39,6 +41,26 @@ def main() -> None:
|
|||
results: dict[str, dict[str, object]] = {}
|
||||
default_project = "runs/yolo26_mlflow_smoke/train" if args.mlflow else "runs/yolo26_smoke"
|
||||
project_dir = (args.project or Path(default_project)).resolve()
|
||||
smoke_batch_id = uuid.uuid4().hex if args.mlflow else ""
|
||||
if args.mlflow:
|
||||
import mlflow
|
||||
|
||||
project_dir.parent.mkdir(parents=True, exist_ok=True)
|
||||
if args.tracking_uri.startswith("sqlite:///"):
|
||||
tracking_db = Path(args.tracking_uri.removeprefix("sqlite:///"))
|
||||
tracking_db.parent.mkdir(parents=True, exist_ok=True)
|
||||
os.environ["YOLO_WEBUI_MLFLOW_RUN_GROUP"] = smoke_batch_id
|
||||
mlflow.set_tracking_uri(args.tracking_uri)
|
||||
mlflow.set_experiment(args.experiment)
|
||||
with mlflow.start_run(run_name=f"smoke-batch-{smoke_batch_id}") as anchor:
|
||||
mlflow.set_tags(
|
||||
{
|
||||
"smoke.anchor": "true",
|
||||
"yolo.run_group": smoke_batch_id,
|
||||
}
|
||||
)
|
||||
print(f"MLflow smoke batch: {smoke_batch_id}", flush=True)
|
||||
|
||||
for task, (dataset, model) in TASKS.items():
|
||||
print(f"\n=== {task}: {model} ===", flush=True)
|
||||
config = TrainingConfig(
|
||||
|
|
@ -83,14 +105,24 @@ def main() -> None:
|
|||
"error": f"{type(exc).__name__}: {exc}",
|
||||
}
|
||||
|
||||
summary_path = project_dir.parent / "smoke_summary.json" if args.mlflow else project_dir / "smoke_summary.json"
|
||||
summary_path = (
|
||||
project_dir.parent / "smoke_summary.json"
|
||||
if args.mlflow
|
||||
else project_dir / "smoke_summary.json"
|
||||
)
|
||||
summary_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
summary = {
|
||||
"smoke_batch_id": smoke_batch_id or None,
|
||||
"tracking_uri": args.tracking_uri if args.mlflow else None,
|
||||
"experiment": args.experiment if args.mlflow else None,
|
||||
"tasks": results,
|
||||
}
|
||||
summary_path.write_text(
|
||||
json.dumps(results, indent=2, ensure_ascii=False) + "\n",
|
||||
json.dumps(summary, indent=2, ensure_ascii=False) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
print(f"\nSummary: {summary_path.resolve()}")
|
||||
print(json.dumps(results, indent=2, ensure_ascii=False))
|
||||
print(json.dumps(summary, indent=2, ensure_ascii=False))
|
||||
if any(result["status"] != "succeeded" for result in results.values()):
|
||||
raise SystemExit(1)
|
||||
|
||||
|
|
|
|||
|
|
@ -3,15 +3,139 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import csv
|
||||
import json
|
||||
import math
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
import mlflow
|
||||
from mlflow.entities import Run
|
||||
from mlflow.tracking import MlflowClient
|
||||
from yolo_webui.mlflow_metrics import PER_CLASS_FIELDS
|
||||
|
||||
|
||||
TASKS = ("detect", "segment", "classify", "pose", "obb")
|
||||
REQUIRED_ARTIFACTS = {"weights/best.pt", "weights/last.pt", "results.csv"}
|
||||
REQUIRED_ARTIFACTS = {
|
||||
"weights/best.pt",
|
||||
"weights/last.pt",
|
||||
"results.csv",
|
||||
"monitoring/task_metrics.csv",
|
||||
}
|
||||
BOX_METRICS = {
|
||||
"metrics/precisionB",
|
||||
"metrics/recallB",
|
||||
"metrics/mAP50B",
|
||||
"metrics/mAP50-95B",
|
||||
}
|
||||
REQUIRED_METRICS_BY_TASK = {
|
||||
"detect": {
|
||||
*BOX_METRICS,
|
||||
"train/box_loss",
|
||||
"train/cls_loss",
|
||||
"train/dfl_loss",
|
||||
"val/box_loss",
|
||||
"val/cls_loss",
|
||||
"val/dfl_loss",
|
||||
"monitor/quality/box/f1_at_optimal_confidence",
|
||||
"monitor/quality/box/map75",
|
||||
},
|
||||
"segment": {
|
||||
*BOX_METRICS,
|
||||
"metrics/precisionM",
|
||||
"metrics/recallM",
|
||||
"metrics/mAP50M",
|
||||
"metrics/mAP50-95M",
|
||||
"train/seg_loss",
|
||||
"train/box_loss",
|
||||
"train/cls_loss",
|
||||
"train/dfl_loss",
|
||||
"val/seg_loss",
|
||||
"val/box_loss",
|
||||
"val/cls_loss",
|
||||
"val/dfl_loss",
|
||||
"monitor/quality/box/f1_at_optimal_confidence",
|
||||
"monitor/quality/mask/f1_at_optimal_confidence",
|
||||
"monitor/quality/mask/map75",
|
||||
},
|
||||
"classify": {
|
||||
"metrics/accuracy_top1",
|
||||
"metrics/accuracy_top5",
|
||||
"train/loss",
|
||||
"val/loss",
|
||||
"monitor/quality/classification/top1_error",
|
||||
"monitor/quality/classification/macro_f1",
|
||||
"monitor/quality/classification/balanced_accuracy",
|
||||
},
|
||||
"pose": {
|
||||
*BOX_METRICS,
|
||||
"metrics/precisionP",
|
||||
"metrics/recallP",
|
||||
"metrics/mAP50P",
|
||||
"metrics/mAP50-95P",
|
||||
"train/pose_loss",
|
||||
"train/box_loss",
|
||||
"train/kobj_loss",
|
||||
"train/cls_loss",
|
||||
"train/dfl_loss",
|
||||
"val/pose_loss",
|
||||
"val/box_loss",
|
||||
"val/kobj_loss",
|
||||
"val/cls_loss",
|
||||
"val/dfl_loss",
|
||||
"monitor/quality/box/f1_at_optimal_confidence",
|
||||
"monitor/quality/keypoints/f1_at_optimal_confidence",
|
||||
"monitor/quality/keypoints/map75",
|
||||
},
|
||||
"obb": {
|
||||
*BOX_METRICS,
|
||||
"train/box_loss",
|
||||
"train/cls_loss",
|
||||
"train/dfl_loss",
|
||||
"train/angle_loss",
|
||||
"val/box_loss",
|
||||
"val/cls_loss",
|
||||
"val/dfl_loss",
|
||||
"val/angle_loss",
|
||||
"monitor/quality/oriented_box/f1_at_optimal_confidence",
|
||||
"monitor/quality/oriented_box/map75",
|
||||
},
|
||||
}
|
||||
COMMON_MONITOR_METRICS = {
|
||||
"monitor/fitness/ultralytics_current",
|
||||
"monitor/fitness/ultralytics_best",
|
||||
"monitor/fitness/task_score_current",
|
||||
"monitor/fitness/task_score_best",
|
||||
"monitor/loss/train_total",
|
||||
"monitor/loss/validation_total",
|
||||
"monitor/loss/generalization_gap",
|
||||
"monitor/loss/validation_to_train_ratio",
|
||||
"monitor/optimization/learning_rate_mean",
|
||||
"monitor/performance/epoch_seconds",
|
||||
"monitor/performance/validation_preprocess_ms_per_image",
|
||||
"monitor/performance/validation_inference_ms_per_image",
|
||||
"monitor/performance/validation_loss_ms_per_image",
|
||||
"monitor/performance/validation_postprocess_ms_per_image",
|
||||
}
|
||||
HISTORY_METRIC_BY_TASK = {
|
||||
"detect": "monitor/quality/box/f1_at_optimal_confidence",
|
||||
"segment": "monitor/quality/mask/f1_at_optimal_confidence",
|
||||
"classify": "monitor/quality/classification/macro_f1",
|
||||
"pose": "monitor/quality/keypoints/f1_at_optimal_confidence",
|
||||
"obb": "monitor/quality/oriented_box/f1_at_optimal_confidence",
|
||||
}
|
||||
EXPECTED_COMPONENTS = {
|
||||
"detect": {"box"},
|
||||
"segment": {"box", "mask"},
|
||||
"classify": {"classification"},
|
||||
"pose": {"box", "keypoints"},
|
||||
"obb": {"oriented_box"},
|
||||
}
|
||||
|
||||
|
||||
def require(condition: object, message: str) -> None:
|
||||
if not condition:
|
||||
raise AssertionError(message)
|
||||
|
||||
|
||||
def artifact_paths(client: MlflowClient, run_id: str, path: str = "") -> set[str]:
|
||||
|
|
@ -24,12 +148,74 @@ def artifact_paths(client: MlflowClient, run_id: str, path: str = "") -> set[str
|
|||
return result
|
||||
|
||||
|
||||
def latest_task_run(runs: list[Run], task: str) -> Run:
|
||||
def latest_task_run(runs: list[Run], task: str, run_group: str) -> Run:
|
||||
expected_name = f"{task}-smoke"
|
||||
for run in runs:
|
||||
if run.data.tags.get("mlflow.runName") == expected_name:
|
||||
if (
|
||||
run.data.tags.get("mlflow.runName") == expected_name
|
||||
and run.data.tags.get("yolo.run_group") == run_group
|
||||
):
|
||||
return run
|
||||
raise AssertionError(f"MLflow run not found: {expected_name}")
|
||||
raise AssertionError(
|
||||
f"MLflow run not found: {expected_name}, yolo.run_group={run_group}"
|
||||
)
|
||||
|
||||
|
||||
def latest_smoke_run_group(runs: list[Run]) -> str:
|
||||
for run in runs:
|
||||
if (
|
||||
run.data.tags.get("smoke.anchor") == "true"
|
||||
and run.data.tags.get("yolo.run_group")
|
||||
):
|
||||
return run.data.tags["yolo.run_group"]
|
||||
raise AssertionError("No MLflow smoke batch with yolo.run_group tag found")
|
||||
|
||||
|
||||
def verify_per_class_artifact(
|
||||
client: MlflowClient,
|
||||
run_id: str,
|
||||
task: str,
|
||||
) -> int:
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
downloaded = Path(
|
||||
client.download_artifacts(
|
||||
run_id,
|
||||
"monitoring/task_metrics.csv",
|
||||
dst_path=directory,
|
||||
)
|
||||
)
|
||||
with downloaded.open(encoding="utf-8") as stream:
|
||||
reader = csv.DictReader(stream)
|
||||
rows = list(reader)
|
||||
require(
|
||||
tuple(reader.fieldnames or ()) == PER_CLASS_FIELDS,
|
||||
f"Invalid task_metrics.csv header for {task}: {reader.fieldnames}",
|
||||
)
|
||||
|
||||
require(rows, f"Empty task_metrics.csv for {task}")
|
||||
components = {row["component"] for row in rows}
|
||||
require(
|
||||
components == EXPECTED_COMPONENTS[task],
|
||||
f"Invalid task_metrics.csv components for {task}: {sorted(components)}",
|
||||
)
|
||||
for row in rows:
|
||||
require(row["task"] == task, f"Invalid task in task_metrics.csv: {row}")
|
||||
require(row["class_id"].isdigit(), f"Invalid class_id in task_metrics.csv: {row}")
|
||||
numeric_fields = ("support", "precision", "recall", "f1")
|
||||
if task != "classify":
|
||||
numeric_fields += ("map50", "map50_95")
|
||||
for field in numeric_fields:
|
||||
try:
|
||||
value = float(row[field])
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise AssertionError(
|
||||
f"Invalid {field} in task_metrics.csv for {task}: {row}"
|
||||
) from exc
|
||||
require(
|
||||
math.isfinite(value),
|
||||
f"Non-finite {field} in task_metrics.csv for {task}: {row}",
|
||||
)
|
||||
return len(rows)
|
||||
|
||||
|
||||
def main() -> None:
|
||||
|
|
@ -39,6 +225,10 @@ def main() -> None:
|
|||
default="sqlite:///runs/yolo26_mlflow_smoke/mlflow.db",
|
||||
)
|
||||
parser.add_argument("--experiment", default="yolo26-mlflow-smoke")
|
||||
parser.add_argument(
|
||||
"--run-group",
|
||||
help="Verify this yolo.run_group tag; defaults to the newest smoke batch.",
|
||||
)
|
||||
args = parser.parse_args()
|
||||
|
||||
mlflow.set_tracking_uri(args.tracking_uri)
|
||||
|
|
@ -51,27 +241,75 @@ def main() -> None:
|
|||
[experiment.experiment_id],
|
||||
order_by=["start_time DESC"],
|
||||
)
|
||||
run_group = args.run_group
|
||||
if not run_group:
|
||||
run_group = latest_smoke_run_group(runs)
|
||||
|
||||
summary: dict[str, object] = {
|
||||
"tracking_uri": args.tracking_uri,
|
||||
"experiment_id": experiment.experiment_id,
|
||||
"artifact_location": experiment.artifact_location,
|
||||
"run_group": run_group,
|
||||
"tasks": {},
|
||||
}
|
||||
task_summary: dict[str, object] = summary["tasks"] # type: ignore[assignment]
|
||||
|
||||
for task in TASKS:
|
||||
run = latest_task_run(runs, task)
|
||||
run = latest_task_run(runs, task, run_group)
|
||||
artifacts = artifact_paths(client, run.info.run_id)
|
||||
missing = REQUIRED_ARTIFACTS - artifacts
|
||||
assert run.info.status == "FINISHED", (task, run.info.status)
|
||||
assert run.data.params, f"No parameters logged for {task}"
|
||||
assert run.data.metrics, f"No metrics logged for {task}"
|
||||
assert not missing, f"Missing artifacts for {task}: {sorted(missing)}"
|
||||
required_metrics = COMMON_MONITOR_METRICS | REQUIRED_METRICS_BY_TASK[task]
|
||||
missing_metrics = required_metrics - run.data.metrics.keys()
|
||||
require(
|
||||
run.info.status == "FINISHED",
|
||||
f"Unexpected run status for {task}: {run.info.status}",
|
||||
)
|
||||
require(run.data.params, f"No parameters logged for {task}")
|
||||
require(
|
||||
not missing_metrics,
|
||||
f"Missing metrics for {task}: {sorted(missing_metrics)}",
|
||||
)
|
||||
invalid_metrics = {
|
||||
key: run.data.metrics[key]
|
||||
for key in required_metrics
|
||||
if not math.isfinite(run.data.metrics[key])
|
||||
}
|
||||
require(
|
||||
not invalid_metrics,
|
||||
f"Non-finite metrics for {task}: {invalid_metrics}",
|
||||
)
|
||||
require(not missing, f"Missing artifacts for {task}: {sorted(missing)}")
|
||||
for metric_key in required_metrics:
|
||||
require(
|
||||
client.get_metric_history(run.info.run_id, metric_key),
|
||||
f"No metric history for {task}: {metric_key}",
|
||||
)
|
||||
history_key = HISTORY_METRIC_BY_TASK[task]
|
||||
history = client.get_metric_history(run.info.run_id, history_key)
|
||||
require(
|
||||
all(math.isfinite(point.value) for point in history),
|
||||
f"Non-finite metric history for {task}: {history_key} {history}",
|
||||
)
|
||||
history_steps = [point.step for point in history]
|
||||
require(
|
||||
history_steps == [0, 1],
|
||||
f"Unexpected metric steps for {task}: {history_key} {history_steps}",
|
||||
)
|
||||
require(run.data.tags.get("yolo.task") == task, f"Missing task tag for {task}")
|
||||
require(
|
||||
run.data.tags.get("monitoring.schema_version") == "1",
|
||||
f"Missing monitoring schema tag for {task}",
|
||||
)
|
||||
per_class_rows = verify_per_class_artifact(client, run.info.run_id, task)
|
||||
task_summary[task] = {
|
||||
"run_id": run.info.run_id,
|
||||
"status": run.info.status,
|
||||
"parameters": len(run.data.params),
|
||||
"metrics": len(run.data.metrics),
|
||||
"required_metrics": sorted(required_metrics),
|
||||
"history_metric": history_key,
|
||||
"history_steps": history_steps,
|
||||
"per_class_rows": per_class_rows,
|
||||
"artifact_uri": run.info.artifact_uri,
|
||||
"required_artifacts": sorted(REQUIRED_ARTIFACTS),
|
||||
}
|
||||
|
|
|
|||
533
src/yolo_webui/mlflow_metrics.py
Normal file
533
src/yolo_webui/mlflow_metrics.py
Normal file
|
|
@ -0,0 +1,533 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import csv
|
||||
import logging
|
||||
import math
|
||||
import os
|
||||
from collections.abc import Iterable, Mapping, MutableMapping
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from .config import YoloTask
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MONITORING_SCHEMA_VERSION = "1"
|
||||
|
||||
# Public component names intentionally do not mirror Ultralytics' one-letter
|
||||
# suffixes. In particular, OBB uses "B" upstream even though its boxes are rotated.
|
||||
TASK_COMPONENTS: dict[YoloTask, tuple[tuple[str, str], ...]] = {
|
||||
"detect": (("box", "box"),),
|
||||
"segment": (("box", "box"), ("mask", "seg")),
|
||||
"classify": (),
|
||||
"pose": (("box", "box"), ("keypoints", "pose")),
|
||||
"obb": (("oriented_box", "box"),),
|
||||
}
|
||||
|
||||
PER_CLASS_FIELDS = (
|
||||
"task",
|
||||
"component",
|
||||
"class_id",
|
||||
"class_name",
|
||||
"support",
|
||||
"precision",
|
||||
"recall",
|
||||
"f1",
|
||||
"map50",
|
||||
"map50_95",
|
||||
)
|
||||
|
||||
|
||||
def _finite_float(value: Any) -> float | None:
|
||||
"""Return a finite Python float for scalar-like values."""
|
||||
if value is None or isinstance(value, (str, bytes, bool)):
|
||||
return None
|
||||
for method in ("detach", "cpu"):
|
||||
operation = getattr(value, method, None)
|
||||
if callable(operation):
|
||||
try:
|
||||
value = operation()
|
||||
except Exception:
|
||||
return None
|
||||
item = getattr(value, "item", None)
|
||||
if callable(item):
|
||||
try:
|
||||
value = item()
|
||||
except Exception:
|
||||
pass
|
||||
try:
|
||||
number = float(value)
|
||||
except (TypeError, ValueError, OverflowError):
|
||||
return None
|
||||
return number if math.isfinite(number) else None
|
||||
|
||||
|
||||
def _finite_values(value: Any) -> list[float]:
|
||||
"""Flatten array-like values while dropping non-finite entries."""
|
||||
if value is None or isinstance(value, (str, bytes, bool)):
|
||||
return []
|
||||
for method in ("detach", "cpu"):
|
||||
operation = getattr(value, method, None)
|
||||
if callable(operation):
|
||||
try:
|
||||
value = operation()
|
||||
except Exception:
|
||||
return []
|
||||
tolist = getattr(value, "tolist", None)
|
||||
if callable(tolist):
|
||||
try:
|
||||
value = tolist()
|
||||
except Exception:
|
||||
return []
|
||||
if isinstance(value, Mapping):
|
||||
source: Iterable[Any] = value.values()
|
||||
elif isinstance(value, Iterable):
|
||||
source = value
|
||||
else:
|
||||
number = _finite_float(value)
|
||||
return [] if number is None else [number]
|
||||
|
||||
result: list[float] = []
|
||||
for item in source:
|
||||
result.extend(_finite_values(item))
|
||||
return result
|
||||
|
||||
|
||||
def _aligned_values(value: Any) -> list[float | None]:
|
||||
"""Convert a one-dimensional array without shifting non-finite positions."""
|
||||
if value is None or isinstance(value, (str, bytes, bool)):
|
||||
return []
|
||||
for method in ("detach", "cpu"):
|
||||
operation = getattr(value, method, None)
|
||||
if callable(operation):
|
||||
try:
|
||||
value = operation()
|
||||
except Exception:
|
||||
return []
|
||||
tolist = getattr(value, "tolist", None)
|
||||
if callable(tolist):
|
||||
try:
|
||||
value = tolist()
|
||||
except Exception:
|
||||
return []
|
||||
if isinstance(value, Mapping):
|
||||
source: Iterable[Any] = value.values()
|
||||
elif isinstance(value, Iterable):
|
||||
source = value
|
||||
else:
|
||||
return [_finite_float(value)]
|
||||
return [_finite_float(item) for item in source]
|
||||
|
||||
|
||||
def _attribute(obj: Any, name: str) -> Any:
|
||||
if obj is None:
|
||||
return None
|
||||
value = getattr(obj, name, None)
|
||||
if callable(value):
|
||||
try:
|
||||
return value()
|
||||
except Exception:
|
||||
return None
|
||||
return value
|
||||
|
||||
|
||||
def _add_metric(metrics: dict[str, float], key: str, value: Any) -> None:
|
||||
number = _finite_float(value)
|
||||
if number is not None:
|
||||
metrics[key] = number
|
||||
|
||||
|
||||
def _mean(values: list[float]) -> float | None:
|
||||
return sum(values) / len(values) if values else None
|
||||
|
||||
|
||||
def _harmonic_mean(precision: float | None, recall: float | None) -> float | None:
|
||||
if precision is None or recall is None:
|
||||
return None
|
||||
denominator = precision + recall
|
||||
return 0.0 if denominator == 0 else 2 * precision * recall / denominator
|
||||
|
||||
|
||||
def _component_metrics(name: str, component: Any) -> dict[str, float]:
|
||||
prefix = f"monitor/quality/{name}"
|
||||
result: dict[str, float] = {}
|
||||
precision = _finite_float(_attribute(component, "mp"))
|
||||
recall = _finite_float(_attribute(component, "mr"))
|
||||
f1_values = _finite_values(_attribute(component, "f1"))
|
||||
ap_values = _finite_values(_attribute(component, "ap"))
|
||||
mean_f1 = _mean(f1_values)
|
||||
|
||||
_add_metric(result, f"{prefix}/precision", precision)
|
||||
_add_metric(result, f"{prefix}/recall", recall)
|
||||
_add_metric(
|
||||
result,
|
||||
f"{prefix}/f1_at_optimal_confidence",
|
||||
mean_f1 if mean_f1 is not None else _harmonic_mean(precision, recall),
|
||||
)
|
||||
_add_metric(result, f"{prefix}/map50", _attribute(component, "map50"))
|
||||
_add_metric(result, f"{prefix}/map75", _attribute(component, "map75"))
|
||||
_add_metric(result, f"{prefix}/map50_95", _attribute(component, "map"))
|
||||
_add_metric(
|
||||
result,
|
||||
f"{prefix}/worst_class_f1_at_optimal_confidence",
|
||||
min(f1_values) if f1_values else None,
|
||||
)
|
||||
_add_metric(result, f"{prefix}/worst_class_map50_95", min(ap_values) if ap_values else None)
|
||||
_add_metric(result, f"{prefix}/classes_evaluated", len(ap_values) or len(f1_values))
|
||||
return result
|
||||
|
||||
|
||||
def _matrix_rows(matrix: Any) -> list[list[float]]:
|
||||
raw_rows = _attribute(matrix, "tolist")
|
||||
if raw_rows is None:
|
||||
raw_rows = matrix
|
||||
if not isinstance(raw_rows, Iterable) or isinstance(raw_rows, (str, bytes)):
|
||||
return []
|
||||
|
||||
rows: list[list[float]] = []
|
||||
for row in raw_rows:
|
||||
if not isinstance(row, Iterable) or isinstance(row, (str, bytes)):
|
||||
return []
|
||||
converted: list[float] = []
|
||||
for value in row:
|
||||
number = _finite_float(value)
|
||||
converted.append(0.0 if number is None else max(0.0, number))
|
||||
rows.append(converted)
|
||||
size = len(rows)
|
||||
return rows if size and all(len(row) == size for row in rows) else []
|
||||
|
||||
|
||||
def _classification_statistics(metric_set: Any) -> tuple[dict[str, float], list[dict[str, Any]]]:
|
||||
prefix = "monitor/quality/classification"
|
||||
result: dict[str, float] = {}
|
||||
top1 = _finite_float(_attribute(metric_set, "top1"))
|
||||
top5 = _finite_float(_attribute(metric_set, "top5"))
|
||||
_add_metric(result, f"{prefix}/top1_accuracy", top1)
|
||||
_add_metric(result, f"{prefix}/top5_accuracy", top5)
|
||||
_add_metric(result, f"{prefix}/top1_error", None if top1 is None else 1.0 - top1)
|
||||
_add_metric(result, f"{prefix}/top5_error", None if top5 is None else 1.0 - top5)
|
||||
|
||||
confusion = _attribute(metric_set, "confusion_matrix")
|
||||
rows = _matrix_rows(_attribute(confusion, "matrix"))
|
||||
names = _attribute(confusion, "names") or {}
|
||||
per_class: list[dict[str, Any]] = []
|
||||
if not rows:
|
||||
return result, per_class
|
||||
|
||||
precisions: list[float] = []
|
||||
supported_recalls: list[float] = []
|
||||
f1_scores: list[float] = []
|
||||
supports: list[float] = []
|
||||
for class_id in range(len(rows)):
|
||||
true_positive = rows[class_id][class_id]
|
||||
predicted = sum(rows[class_id])
|
||||
support = sum(row[class_id] for row in rows)
|
||||
if support <= 0 and predicted <= 0:
|
||||
continue
|
||||
precision = true_positive / predicted if predicted else 0.0
|
||||
recall = true_positive / support if support else 0.0
|
||||
f1 = _harmonic_mean(precision, recall) or 0.0
|
||||
precisions.append(precision)
|
||||
f1_scores.append(f1)
|
||||
supports.append(support)
|
||||
if support > 0:
|
||||
supported_recalls.append(recall)
|
||||
per_class.append(
|
||||
{
|
||||
"task": "classify",
|
||||
"component": "classification",
|
||||
"class_id": class_id,
|
||||
"class_name": _class_name(names, class_id),
|
||||
"support": support,
|
||||
"precision": precision,
|
||||
"recall": recall,
|
||||
"f1": f1,
|
||||
"map50": "",
|
||||
"map50_95": "",
|
||||
}
|
||||
)
|
||||
|
||||
total_support = sum(supports)
|
||||
weighted_f1 = (
|
||||
sum(score * support for score, support in zip(f1_scores, supports)) / total_support
|
||||
if total_support
|
||||
else None
|
||||
)
|
||||
_add_metric(result, f"{prefix}/macro_precision", _mean(precisions))
|
||||
_add_metric(result, f"{prefix}/macro_recall", _mean(supported_recalls))
|
||||
_add_metric(result, f"{prefix}/macro_f1", _mean(f1_scores))
|
||||
_add_metric(result, f"{prefix}/weighted_f1", weighted_f1)
|
||||
_add_metric(result, f"{prefix}/balanced_accuracy", _mean(supported_recalls))
|
||||
_add_metric(
|
||||
result,
|
||||
f"{prefix}/worst_class_recall",
|
||||
min(supported_recalls) if supported_recalls else None,
|
||||
)
|
||||
_add_metric(result, f"{prefix}/classes_evaluated", len(per_class))
|
||||
return result, per_class
|
||||
|
||||
|
||||
def collect_monitoring_metrics(task: YoloTask, trainer: Any) -> dict[str, float]:
|
||||
"""Collect stable, task-aware metrics from an Ultralytics trainer."""
|
||||
result: dict[str, float] = {}
|
||||
trainer_metrics = getattr(trainer, "metrics", {}) or {}
|
||||
validator = getattr(trainer, "validator", None)
|
||||
metric_set = getattr(validator, "metrics", None)
|
||||
final_validation = (
|
||||
validator is not None and getattr(validator, "training", True) is False
|
||||
)
|
||||
|
||||
current_fitness = _finite_float(_attribute(metric_set, "fitness"))
|
||||
if current_fitness is None:
|
||||
current_fitness = _finite_float(getattr(trainer, "fitness", None))
|
||||
best_fitness = _finite_float(getattr(trainer, "best_fitness", None))
|
||||
_add_metric(result, "monitor/fitness/ultralytics_current", current_fitness)
|
||||
_add_metric(result, "monitor/fitness/ultralytics_best", best_fitness)
|
||||
fitness_scale = max(1, len(TASK_COMPONENTS[task]))
|
||||
_add_metric(
|
||||
result,
|
||||
"monitor/fitness/task_score_current",
|
||||
None if current_fitness is None else current_fitness / fitness_scale,
|
||||
)
|
||||
_add_metric(
|
||||
result,
|
||||
"monitor/fitness/task_score_best",
|
||||
None if best_fitness is None else best_fitness / fitness_scale,
|
||||
)
|
||||
|
||||
if not final_validation:
|
||||
train_total: float | None = None
|
||||
validation_total: float | None = None
|
||||
label_losses = getattr(trainer, "label_loss_items", None)
|
||||
if callable(label_losses):
|
||||
try:
|
||||
train_losses = label_losses(getattr(trainer, "tloss", None), prefix="train")
|
||||
except Exception:
|
||||
train_losses = {}
|
||||
if isinstance(train_losses, Mapping):
|
||||
values = [
|
||||
number
|
||||
for value in train_losses.values()
|
||||
if (number := _finite_float(value)) is not None
|
||||
]
|
||||
train_total = sum(values) if values else None
|
||||
_add_metric(result, "monitor/loss/train_total", train_total)
|
||||
|
||||
if isinstance(trainer_metrics, Mapping):
|
||||
validation_loss_keys: set[str] = set()
|
||||
if callable(label_losses):
|
||||
try:
|
||||
expected_losses = label_losses(None, prefix="val")
|
||||
except Exception:
|
||||
expected_losses = ()
|
||||
if isinstance(expected_losses, Mapping):
|
||||
validation_loss_keys.update(map(str, expected_losses))
|
||||
elif isinstance(expected_losses, Iterable) and not isinstance(
|
||||
expected_losses,
|
||||
(str, bytes),
|
||||
):
|
||||
validation_loss_keys.update(map(str, expected_losses))
|
||||
if not validation_loss_keys:
|
||||
validation_loss_keys = {
|
||||
str(key)
|
||||
for key in trainer_metrics
|
||||
if str(key) == "val/loss"
|
||||
or (
|
||||
str(key).startswith("val/")
|
||||
and str(key).endswith("_loss")
|
||||
)
|
||||
}
|
||||
val_losses = [
|
||||
number
|
||||
for key, value in trainer_metrics.items()
|
||||
if str(key) in validation_loss_keys
|
||||
and (number := _finite_float(value)) is not None
|
||||
]
|
||||
validation_total = sum(val_losses) if val_losses else None
|
||||
_add_metric(result, "monitor/loss/validation_total", validation_total)
|
||||
|
||||
if train_total is not None and validation_total is not None:
|
||||
_add_metric(
|
||||
result,
|
||||
"monitor/loss/generalization_gap",
|
||||
validation_total - train_total,
|
||||
)
|
||||
if train_total > 0:
|
||||
_add_metric(
|
||||
result,
|
||||
"monitor/loss/validation_to_train_ratio",
|
||||
validation_total / train_total,
|
||||
)
|
||||
|
||||
learning_rates = _finite_values(getattr(trainer, "lr", None))
|
||||
_add_metric(result, "monitor/optimization/learning_rate_mean", _mean(learning_rates))
|
||||
_add_metric(
|
||||
result,
|
||||
"monitor/performance/epoch_seconds",
|
||||
getattr(trainer, "epoch_time", None),
|
||||
)
|
||||
|
||||
speed = _attribute(metric_set, "speed")
|
||||
if isinstance(speed, Mapping):
|
||||
for stage in ("preprocess", "inference", "loss", "postprocess"):
|
||||
_add_metric(
|
||||
result,
|
||||
f"monitor/performance/validation_{stage}_ms_per_image",
|
||||
speed.get(stage),
|
||||
)
|
||||
|
||||
if task == "classify":
|
||||
classification, _ = _classification_statistics(metric_set)
|
||||
result.update(classification)
|
||||
else:
|
||||
for public_name, attribute_name in TASK_COMPONENTS[task]:
|
||||
component = _attribute(metric_set, attribute_name)
|
||||
if component is not None:
|
||||
result.update(_component_metrics(public_name, component))
|
||||
return result
|
||||
|
||||
|
||||
def _class_name(names: Any, class_id: int) -> str:
|
||||
if isinstance(names, Mapping):
|
||||
return str(names.get(class_id, names.get(str(class_id), class_id)))
|
||||
if isinstance(names, (list, tuple)) and 0 <= class_id < len(names):
|
||||
return str(names[class_id])
|
||||
return str(class_id)
|
||||
|
||||
|
||||
def _value_at(values: list[float | None], index: int) -> float | str:
|
||||
if index >= len(values) or values[index] is None:
|
||||
return ""
|
||||
return values[index]
|
||||
|
||||
|
||||
def _component_rows(
|
||||
task: YoloTask,
|
||||
public_name: str,
|
||||
component: Any,
|
||||
metric_set: Any,
|
||||
) -> list[dict[str, Any]]:
|
||||
raw_class_indices = _aligned_values(_attribute(component, "ap_class_index"))
|
||||
precision = _aligned_values(_attribute(component, "p"))
|
||||
recall = _aligned_values(_attribute(component, "r"))
|
||||
f1 = _aligned_values(_attribute(component, "f1"))
|
||||
map50 = _aligned_values(_attribute(component, "ap50"))
|
||||
map50_95 = _aligned_values(_attribute(component, "ap"))
|
||||
count = max(
|
||||
map(len, (raw_class_indices, precision, recall, f1, map50, map50_95)),
|
||||
default=0,
|
||||
)
|
||||
class_indices = [
|
||||
index if value is None else int(value)
|
||||
for index, value in enumerate(raw_class_indices)
|
||||
]
|
||||
if not class_indices:
|
||||
class_indices = list(range(count))
|
||||
|
||||
names = _attribute(metric_set, "names") or {}
|
||||
support_by_class = _aligned_values(_attribute(metric_set, "nt_per_class"))
|
||||
rows: list[dict[str, Any]] = []
|
||||
for index in range(min(count, len(class_indices))):
|
||||
class_id = class_indices[index]
|
||||
rows.append(
|
||||
{
|
||||
"task": task,
|
||||
"component": public_name,
|
||||
"class_id": class_id,
|
||||
"class_name": _class_name(names, class_id),
|
||||
"support": _value_at(support_by_class, class_id),
|
||||
"precision": _value_at(precision, index),
|
||||
"recall": _value_at(recall, index),
|
||||
"f1": _value_at(f1, index),
|
||||
"map50": _value_at(map50, index),
|
||||
"map50_95": _value_at(map50_95, index),
|
||||
}
|
||||
)
|
||||
return rows
|
||||
|
||||
|
||||
def collect_per_class_metrics(task: YoloTask, trainer: Any) -> list[dict[str, Any]]:
|
||||
"""Build final per-class diagnostics for the MLflow CSV artifact."""
|
||||
validator = getattr(trainer, "validator", None)
|
||||
metric_set = getattr(validator, "metrics", None)
|
||||
if metric_set is None:
|
||||
return []
|
||||
if task == "classify":
|
||||
_, rows = _classification_statistics(metric_set)
|
||||
return rows
|
||||
|
||||
result: list[dict[str, Any]] = []
|
||||
for public_name, attribute_name in TASK_COMPONENTS[task]:
|
||||
component = _attribute(metric_set, attribute_name)
|
||||
if component is not None:
|
||||
result.extend(_component_rows(task, public_name, component, metric_set))
|
||||
return result
|
||||
|
||||
|
||||
class TaskMetricsMonitor:
|
||||
"""Enrich Ultralytics metrics before its built-in MLflow callback runs."""
|
||||
|
||||
def __init__(self, task: YoloTask, *, mlflow_enabled: bool) -> None:
|
||||
self.task = task
|
||||
self.mlflow_enabled = mlflow_enabled
|
||||
|
||||
def on_train_start(self, trainer: Any) -> None:
|
||||
"""Attach searchable task/schema tags without ever starting a second run."""
|
||||
if not self.mlflow_enabled or not getattr(trainer, "_mlflow_active", False):
|
||||
return
|
||||
try:
|
||||
import mlflow
|
||||
|
||||
if mlflow.active_run() is not None:
|
||||
tags = {
|
||||
"monitoring.schema_version": MONITORING_SCHEMA_VERSION,
|
||||
"yolo.task": self.task,
|
||||
}
|
||||
run_group = os.environ.get("YOLO_WEBUI_MLFLOW_RUN_GROUP", "").strip()
|
||||
if run_group:
|
||||
tags["yolo.run_group"] = run_group
|
||||
mlflow.set_tags(tags)
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Не удалось записать теги мониторинга "
|
||||
"в MLflow: %s",
|
||||
exc,
|
||||
)
|
||||
|
||||
def on_fit_epoch_end(self, trainer: Any) -> None:
|
||||
"""Add derived metrics to trainer.metrics for the current validation epoch."""
|
||||
if not self.mlflow_enabled:
|
||||
return
|
||||
metrics = getattr(trainer, "metrics", None)
|
||||
if not isinstance(metrics, MutableMapping):
|
||||
return
|
||||
try:
|
||||
metrics.update(collect_monitoring_metrics(self.task, trainer))
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Не удалось собрать task-specific метрики: %s",
|
||||
exc,
|
||||
)
|
||||
|
||||
def on_train_end(self, trainer: Any) -> None:
|
||||
"""Write and upload final per-class diagnostics to the active MLflow run."""
|
||||
if not self.mlflow_enabled:
|
||||
return
|
||||
try:
|
||||
rows = collect_per_class_metrics(self.task, trainer)
|
||||
output = Path(trainer.save_dir) / "monitoring" / "task_metrics.csv"
|
||||
output.parent.mkdir(parents=True, exist_ok=True)
|
||||
with output.open("w", newline="", encoding="utf-8") as stream:
|
||||
writer = csv.DictWriter(stream, fieldnames=PER_CLASS_FIELDS)
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
if getattr(trainer, "_mlflow_active", False):
|
||||
import mlflow
|
||||
|
||||
if mlflow.active_run() is not None:
|
||||
mlflow.log_artifact(str(output), artifact_path="monitoring")
|
||||
except Exception as exc:
|
||||
logger.warning(
|
||||
"Не удалось сохранить per-class метрики: %s",
|
||||
exc,
|
||||
)
|
||||
|
|
@ -189,6 +189,7 @@ class TrainingRunner:
|
|||
|
||||
on_event(TrainingEvent("info", "Загрузка Ultralytics и подготовка модели…"))
|
||||
from ultralytics import YOLO, settings
|
||||
from .ultralytics_trainers import trainer_for_task
|
||||
|
||||
settings.update({"mlflow": config.mlflow.enabled})
|
||||
|
||||
|
|
@ -198,11 +199,11 @@ class TrainingRunner:
|
|||
self._model = model
|
||||
|
||||
model.add_callback("on_train_start", self._on_train_start(on_event))
|
||||
model.add_callback("on_train_epoch_end", self._on_epoch_end(on_event))
|
||||
model.add_callback("on_fit_epoch_end", self._on_epoch_end(on_event))
|
||||
model.add_callback("on_train_end", self._on_train_end(on_event))
|
||||
|
||||
try:
|
||||
model.train(**train_args)
|
||||
model.train(trainer=trainer_for_task(config.task), **train_args)
|
||||
trainer = getattr(model, "trainer", None)
|
||||
save_dir = getattr(trainer, "save_dir", None)
|
||||
return Path(save_dir) if save_dir else None
|
||||
|
|
@ -221,6 +222,9 @@ class TrainingRunner:
|
|||
|
||||
def _on_epoch_end(self, on_event: EventHandler) -> Callable[[Any], None]:
|
||||
def callback(trainer: Any) -> None:
|
||||
validator = getattr(trainer, "validator", None)
|
||||
if validator is not None and getattr(validator, "training", True) is False:
|
||||
return # final best-checkpoint validation is not an additional train epoch
|
||||
epoch = int(getattr(trainer, "epoch", 0)) + 1
|
||||
total = int(getattr(getattr(trainer, "args", None), "epochs", 0))
|
||||
metrics = getattr(trainer, "metrics", {}) or {}
|
||||
|
|
|
|||
72
src/yolo_webui/ultralytics_trainers.py
Normal file
72
src/yolo_webui/ultralytics_trainers.py
Normal file
|
|
@ -0,0 +1,72 @@
|
|||
from __future__ import annotations
|
||||
|
||||
from typing import Any, ClassVar
|
||||
|
||||
from ultralytics.models.yolo.classify.train import ClassificationTrainer
|
||||
from ultralytics.models.yolo.detect.train import DetectionTrainer
|
||||
from ultralytics.models.yolo.obb.train import OBBTrainer
|
||||
from ultralytics.models.yolo.pose.train import PoseTrainer
|
||||
from ultralytics.models.yolo.segment.train import SegmentationTrainer
|
||||
from ultralytics.utils import SETTINGS
|
||||
|
||||
from .config import YoloTask
|
||||
from .mlflow_metrics import TaskMetricsMonitor
|
||||
|
||||
|
||||
class _TaskMetricsTrainerMixin:
|
||||
"""Install monitoring inside the trainer so Ultralytics DDP keeps it."""
|
||||
|
||||
monitoring_task: ClassVar[YoloTask]
|
||||
|
||||
def __init__(self, *args: Any, **kwargs: Any) -> None:
|
||||
super().__init__(*args, **kwargs)
|
||||
self._install_task_metrics_monitor()
|
||||
|
||||
def _install_task_metrics_monitor(self) -> None:
|
||||
monitor = TaskMetricsMonitor(
|
||||
self.monitoring_task,
|
||||
mlflow_enabled=SETTINGS["mlflow"] is True,
|
||||
)
|
||||
self._task_metrics_monitor = monitor
|
||||
callbacks = (
|
||||
("on_train_start", monitor.on_train_start),
|
||||
("on_fit_epoch_end", monitor.on_fit_epoch_end),
|
||||
("on_train_end", monitor.on_train_end),
|
||||
)
|
||||
for event, callback in callbacks:
|
||||
# BaseTrainer adds integrations in __init__. Prepending ensures our
|
||||
# metrics and CSV exist before Ultralytics logs them to MLflow.
|
||||
self.callbacks.setdefault(event, []).insert(0, callback)
|
||||
|
||||
|
||||
class MonitoredDetectionTrainer(_TaskMetricsTrainerMixin, DetectionTrainer):
|
||||
monitoring_task = "detect"
|
||||
|
||||
|
||||
class MonitoredSegmentationTrainer(_TaskMetricsTrainerMixin, SegmentationTrainer):
|
||||
monitoring_task = "segment"
|
||||
|
||||
|
||||
class MonitoredClassificationTrainer(_TaskMetricsTrainerMixin, ClassificationTrainer):
|
||||
monitoring_task = "classify"
|
||||
|
||||
|
||||
class MonitoredPoseTrainer(_TaskMetricsTrainerMixin, PoseTrainer):
|
||||
monitoring_task = "pose"
|
||||
|
||||
|
||||
class MonitoredOBBTrainer(_TaskMetricsTrainerMixin, OBBTrainer):
|
||||
monitoring_task = "obb"
|
||||
|
||||
|
||||
MONITORED_TRAINERS = {
|
||||
"detect": MonitoredDetectionTrainer,
|
||||
"segment": MonitoredSegmentationTrainer,
|
||||
"classify": MonitoredClassificationTrainer,
|
||||
"pose": MonitoredPoseTrainer,
|
||||
"obb": MonitoredOBBTrainer,
|
||||
}
|
||||
|
||||
|
||||
def trainer_for_task(task: YoloTask) -> type:
|
||||
return MONITORED_TRAINERS[task]
|
||||
423
tests/test_mlflow_metrics.py
Normal file
423
tests/test_mlflow_metrics.py
Normal file
|
|
@ -0,0 +1,423 @@
|
|||
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)
|
||||
43
tests/test_mlflow_smoke_verifier.py
Normal file
43
tests/test_mlflow_smoke_verifier.py
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
from __future__ import annotations
|
||||
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
SPEC = importlib.util.spec_from_file_location(
|
||||
"verify_mlflow_smoke",
|
||||
Path(__file__).parents[1] / "scripts" / "verify_mlflow_smoke.py",
|
||||
)
|
||||
assert SPEC is not None and SPEC.loader is not None
|
||||
verify_mlflow_smoke = importlib.util.module_from_spec(SPEC)
|
||||
SPEC.loader.exec_module(verify_mlflow_smoke)
|
||||
latest_smoke_run_group = verify_mlflow_smoke.latest_smoke_run_group
|
||||
latest_task_run = verify_mlflow_smoke.latest_task_run
|
||||
|
||||
|
||||
def _run(name: str, group: str, *, anchor: bool = False) -> SimpleNamespace:
|
||||
tags = {
|
||||
"mlflow.runName": name,
|
||||
"yolo.run_group": group,
|
||||
}
|
||||
if anchor:
|
||||
tags["smoke.anchor"] = "true"
|
||||
return SimpleNamespace(data=SimpleNamespace(tags=tags))
|
||||
|
||||
|
||||
def test_verifier_selects_one_complete_smoke_batch() -> None:
|
||||
runs = [
|
||||
_run("detect-smoke", "new"),
|
||||
_run("smoke-batch-new", "new", anchor=True),
|
||||
_run("obb-smoke", "old"),
|
||||
_run("smoke-batch-old", "old", anchor=True),
|
||||
]
|
||||
|
||||
group = latest_smoke_run_group(runs)
|
||||
|
||||
assert group == "new"
|
||||
assert latest_task_run(runs, "detect", group).data.tags["yolo.run_group"] == "new"
|
||||
with pytest.raises(AssertionError, match="obb-smoke"):
|
||||
latest_task_run(runs, "obb", group)
|
||||
|
|
@ -43,31 +43,46 @@ def test_runner_wires_yolo_callbacks_and_returns_output(
|
|||
class FakeYOLO:
|
||||
def __init__(self, model: str, task: str) -> None:
|
||||
constructed.append((model, task))
|
||||
self.callbacks: dict[str, Any] = {}
|
||||
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[name] = callback
|
||||
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.callbacks["on_train_start"](self.trainer)
|
||||
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.callbacks["on_train_epoch_end"](self.trainer)
|
||||
self.callbacks["on_train_end"](self.trainer)
|
||||
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",
|
||||
|
|
@ -85,7 +100,24 @@ def test_runner_wires_yolo_callbacks_and_returns_output(
|
|||
assert settings_updates == [{"mlflow": False}]
|
||||
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_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:
|
||||
|
|
|
|||
39
tests/test_ultralytics_trainers.py
Normal file
39
tests/test_ultralytics_trainers.py
Normal file
|
|
@ -0,0 +1,39 @@
|
|||
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
|
||||
Loading…
Add table
Reference in a new issue