65 lines
2.2 KiB
Python
65 lines
2.2 KiB
Python
"""Base model versus new model, on the same val set (REQ-063).
|
|
|
|
Both models are validated against one `data.yaml`, so the numbers differ only
|
|
because the weights differ. Combined with the stable val split in `dataset.py`,
|
|
that is what makes "the model improved" a claim rather than a hope.
|
|
"""
|
|
|
|
from typing import Optional
|
|
|
|
|
|
class EvaluateError(Exception):
|
|
pass
|
|
|
|
|
|
def class_names(weights: str) -> list:
|
|
from ultralytics import YOLO
|
|
|
|
names = YOLO(weights).names
|
|
return [names[key] for key in sorted(names)] if isinstance(names, dict) else list(names)
|
|
|
|
|
|
def validate(weights: str, data_yaml: str, imgsz: int = 640, device=0,
|
|
batch: int = 8) -> dict:
|
|
from ultralytics import YOLO
|
|
|
|
metrics = YOLO(weights).val(
|
|
data=data_yaml, imgsz=imgsz, device=device, batch=batch,
|
|
split="val", plots=False, verbose=False,
|
|
)
|
|
box = metrics.box
|
|
return {
|
|
"map50": round(float(box.map50), 4),
|
|
"map50_95": round(float(box.map), 4),
|
|
"precision": round(float(box.mp), 4),
|
|
"recall": round(float(box.mr), 4),
|
|
}
|
|
|
|
|
|
def compare(base_weights: Optional[str], new_weights: str, data_yaml: str,
|
|
expected_classes: list, imgsz: int = 640, device=0,
|
|
batch: int = 8) -> dict:
|
|
"""Validate both models where that is meaningful, and say so when it is not.
|
|
|
|
A base model whose class list differs from the project's cannot be scored on
|
|
this dataset — its class ids mean something else. Reporting nothing beats
|
|
reporting a number that looks like a regression but is a mismatch.
|
|
"""
|
|
new_metrics = validate(new_weights, data_yaml, imgsz, device, batch)
|
|
|
|
base_metrics = None
|
|
skipped = None
|
|
if not base_weights:
|
|
skipped = "This project has no base model yet — nothing to compare against."
|
|
else:
|
|
try:
|
|
base_metrics = validate(base_weights, data_yaml, imgsz, device, batch)
|
|
except Exception as exc:
|
|
base_metrics, skipped = None, f"Could not evaluate base model: {exc}"
|
|
|
|
delta = None
|
|
if base_metrics:
|
|
delta = {key: round(new_metrics[key] - base_metrics[key], 4) for key in new_metrics}
|
|
|
|
return {"base": base_metrics, "new": new_metrics, "delta": delta, "skipped": skipped}
|