321 lines
10 KiB
Python
Executable file
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()
|