train_utility/scripts/verify_mlflow_smoke.py
2026-08-04 11:14:51 +04:00

321 lines
10 KiB
Python
Executable file

"""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()