"""Verify that every YOLO smoke task was persisted completely in MLflow.""" 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", "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]: result: set[str] = set() for artifact in client.list_artifacts(run_id, path): if artifact.is_dir: result.update(artifact_paths(client, run_id, artifact.path)) else: result.add(artifact.path) return result 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 and run.data.tags.get("yolo.run_group") == run_group ): return run 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: parser = argparse.ArgumentParser() parser.add_argument( "--tracking-uri", 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) client = MlflowClient() experiment = client.get_experiment_by_name(args.experiment) if experiment is None: raise AssertionError(f"MLflow experiment not found: {args.experiment}") runs = client.search_runs( [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_group) artifacts = artifact_paths(client, run.info.run_id) missing = REQUIRED_ARTIFACTS - artifacts 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), } print(json.dumps(summary, indent=2, ensure_ascii=False)) if __name__ == "__main__": main()