Download evaluate.py from AI-MED-AGH/Recruitment-Task-3-Model: direct link, hf CLI and curl.
- Browser
- Download file 16 kB
-
https://huggingface.co/AI-MED-AGH/Recruitment-Task-3-Model/resolve/main/evaluate.py
- Command line
-
hf download hf://AI-MED-AGH/Recruitment-Task-3-Model/evaluate.py
-
curl -L -o evaluate.py https://huggingface.co/AI-MED-AGH/Recruitment-Task-3-Model/resolve/main/evaluate.py
16 kB
| """Frozen-checkpoint evaluation and artifact export for DeepWeeds.""" | |
| from __future__ import annotations | |
| import argparse | |
| import csv | |
| import json | |
| import math | |
| import platform | |
| import sys | |
| import time | |
| from dataclasses import dataclass | |
| from decimal import Decimal | |
| from importlib import metadata | |
| from pathlib import Path | |
| from typing import Any, Sequence | |
| import torch | |
| from safetensors.torch import load_file | |
| from sklearn.metrics import ( | |
| accuracy_score, | |
| balanced_accuracy_score, | |
| f1_score, | |
| precision_recall_fscore_support, | |
| ) | |
| from torch import Tensor, nn | |
| from torch.utils.data import DataLoader | |
| try: # Package import in the organizer workspace. | |
| from .data import build_comparison_transforms, load_fold_splits, make_loader | |
| from .model import SmallDeepWeedsCNN | |
| except ImportError: # Flat-file import after copying into Hub staging. | |
| from data import build_comparison_transforms, load_fold_splits, make_loader | |
| from model import SmallDeepWeedsCNN | |
| PUBLICATION_THRESHOLD = Decimal("0.10") | |
| DEFAULT_CLASS_NAMES = tuple(str(class_id) for class_id in range(9)) | |
| class PredictionResult: | |
| """Aligned outputs from one complete, non-shuffled evaluation pass.""" | |
| filenames: list[str] | |
| true_labels: list[int] | |
| predicted_labels: list[int] | |
| probabilities: list[list[float]] | |
| inference_seconds: float | |
| def _json_ready(value: Any) -> Any: | |
| """Return JSON-compatible builtins while rejecting non-finite floats.""" | |
| if isinstance(value, dict): | |
| return {str(key): _json_ready(item) for key, item in value.items()} | |
| if isinstance(value, (list, tuple)): | |
| return [_json_ready(item) for item in value] | |
| if isinstance(value, Path): | |
| return str(value) | |
| if isinstance(value, bool) or value is None or isinstance(value, (str, int)): | |
| return value | |
| if isinstance(value, float): | |
| if not math.isfinite(value): | |
| raise ValueError(f"JSON artifacts require finite numbers, got {value!r}") | |
| return value | |
| if hasattr(value, "item"): | |
| return _json_ready(value.item()) | |
| raise TypeError(f"Unsupported JSON artifact value: {type(value).__name__}") | |
| def write_json(path: str | Path, payload: dict[str, Any]) -> dict[str, Any]: | |
| """Write a deterministic, strictly finite JSON object.""" | |
| destination = Path(path).resolve() | |
| destination.parent.mkdir(parents=True, exist_ok=True) | |
| ready = _json_ready(payload) | |
| destination.write_text( | |
| json.dumps(ready, indent=2, sort_keys=True, allow_nan=False) + "\n", | |
| encoding="utf-8", | |
| ) | |
| return ready | |
| def compute_classification_metrics( | |
| true_labels: Sequence[int], | |
| predicted_labels: Sequence[int], | |
| class_names: Sequence[str] = DEFAULT_CLASS_NAMES, | |
| ) -> dict[str, Any]: | |
| """Compute the declared multiclass metrics with explicit zero handling.""" | |
| truth = [int(label) for label in true_labels] | |
| predictions = [int(label) for label in predicted_labels] | |
| names = [str(name) for name in class_names] | |
| if not truth: | |
| raise ValueError("At least one labelled prediction is required") | |
| if len(truth) != len(predictions): | |
| raise ValueError("true_labels and predicted_labels must have the same length") | |
| if not names: | |
| raise ValueError("class_names must not be empty") | |
| labels = list(range(len(names))) | |
| if any(label not in labels for label in truth + predictions): | |
| raise ValueError("All labels must index class_names") | |
| precision, recall, per_class_f1, support = precision_recall_fscore_support( | |
| truth, | |
| predictions, | |
| labels=labels, | |
| zero_division=0, | |
| ) | |
| metrics = { | |
| "accuracy": float(accuracy_score(truth, predictions)), | |
| "macro_f1": float( | |
| f1_score( | |
| truth, | |
| predictions, | |
| labels=labels, | |
| average="macro", | |
| zero_division=0, | |
| ) | |
| ), | |
| "balanced_accuracy": float(balanced_accuracy_score(truth, predictions)), | |
| "per_class": [ | |
| { | |
| "class_id": class_id, | |
| "class_name": names[class_id], | |
| "precision": float(precision[class_id]), | |
| "recall": float(recall[class_id]), | |
| "f1": float(per_class_f1[class_id]), | |
| "support": int(support[class_id]), | |
| } | |
| for class_id in labels | |
| ], | |
| } | |
| return _json_ready(metrics) | |
| def _synchronize(device: torch.device) -> None: | |
| if device.type == "cuda": | |
| torch.cuda.synchronize(device) | |
| def predict_model( | |
| model: nn.Module, | |
| loader: DataLoader, | |
| device: torch.device, | |
| ) -> PredictionResult: | |
| """Run one complete inference pass while retaining row alignment.""" | |
| model.to(device) | |
| model.eval() | |
| filenames: list[str] = [] | |
| true_labels: list[int] = [] | |
| predicted_labels: list[int] = [] | |
| probabilities: list[list[float]] = [] | |
| _synchronize(device) | |
| started = time.perf_counter() | |
| with torch.inference_mode(): | |
| for images, labels, batch_filenames in loader: | |
| logits = model(images.to(device, non_blocking=device.type == "cuda")) | |
| if logits.ndim != 2: | |
| raise ValueError(f"Expected two-dimensional logits, got {tuple(logits.shape)}") | |
| batch_probabilities = torch.softmax(logits, dim=1).cpu() | |
| batch_predictions = batch_probabilities.argmax(dim=1) | |
| filenames.extend(str(filename) for filename in batch_filenames) | |
| true_labels.extend(int(label) for label in labels.tolist()) | |
| predicted_labels.extend(int(label) for label in batch_predictions.tolist()) | |
| probabilities.extend( | |
| [float(probability) for probability in row] | |
| for row in batch_probabilities.tolist() | |
| ) | |
| _synchronize(device) | |
| elapsed = time.perf_counter() - started | |
| if not filenames: | |
| raise ValueError("Evaluation loader produced no prediction rows") | |
| if len(set(filenames)) != len(filenames): | |
| raise ValueError("Evaluation filenames must be unique") | |
| return PredictionResult( | |
| filenames=filenames, | |
| true_labels=true_labels, | |
| predicted_labels=predicted_labels, | |
| probabilities=probabilities, | |
| inference_seconds=elapsed, | |
| ) | |
| def write_predictions_csv( | |
| path: str | Path, | |
| filenames: Sequence[str], | |
| true_labels: Sequence[int], | |
| predicted_labels: Sequence[int], | |
| probabilities: Sequence[Sequence[float]], | |
| ) -> Path: | |
| """Write one integrity-checked prediction row per source filename.""" | |
| row_counts = { | |
| len(filenames), len(true_labels), len(predicted_labels), len(probabilities) | |
| } | |
| if len(row_counts) != 1: | |
| raise ValueError("Prediction fields must contain the same number of rows") | |
| if not filenames: | |
| raise ValueError("Prediction export must contain at least one row") | |
| if len(set(str(filename) for filename in filenames)) != len(filenames): | |
| raise ValueError("Prediction filenames must be unique") | |
| class_counts = {len(row) for row in probabilities} | |
| if len(class_counts) != 1 or not class_counts or next(iter(class_counts)) <= 0: | |
| raise ValueError("Every probability row must have the same positive class count") | |
| class_count = next(iter(class_counts)) | |
| for row in probabilities: | |
| if any(not math.isfinite(float(value)) for value in row): | |
| raise ValueError("Prediction probabilities must be finite") | |
| if any(not 0 <= int(label) < class_count for label in true_labels): | |
| raise ValueError("A true label is outside the probability columns") | |
| if any(not 0 <= int(label) < class_count for label in predicted_labels): | |
| raise ValueError("A predicted label is outside the probability columns") | |
| destination = Path(path).resolve() | |
| destination.parent.mkdir(parents=True, exist_ok=True) | |
| probability_fields = [f"probability_{class_id}" for class_id in range(class_count)] | |
| with destination.open("w", newline="", encoding="utf-8") as csv_file: | |
| writer = csv.DictWriter( | |
| csv_file, | |
| fieldnames=["filename", "true_label", "predicted_label", *probability_fields], | |
| ) | |
| writer.writeheader() | |
| for filename, true_label, predicted_label, row in zip( | |
| filenames, true_labels, predicted_labels, probabilities, strict=True | |
| ): | |
| writer.writerow( | |
| { | |
| "filename": str(filename), | |
| "true_label": int(true_label), | |
| "predicted_label": int(predicted_label), | |
| **{ | |
| field: float(probability) | |
| for field, probability in zip(probability_fields, row, strict=True) | |
| }, | |
| } | |
| ) | |
| return destination | |
| def package_environment(device: torch.device) -> dict[str, Any]: | |
| """Capture the runtime facts needed to reproduce an evaluation.""" | |
| def version(distribution: str) -> str: | |
| try: | |
| return metadata.version(distribution) | |
| except metadata.PackageNotFoundError: | |
| return "not-installed" | |
| return { | |
| "python": platform.python_version(), | |
| "platform": platform.platform(), | |
| "device": str(device), | |
| "cuda_available": torch.cuda.is_available(), | |
| "cuda_device_name": ( | |
| torch.cuda.get_device_name(device) if device.type == "cuda" else None | |
| ), | |
| "packages": { | |
| "torch": torch.__version__, | |
| "torchvision": version("torchvision"), | |
| "scikit_learn": version("scikit-learn"), | |
| "safetensors": version("safetensors"), | |
| }, | |
| } | |
| def evaluate_checkpoint( | |
| model: nn.Module, | |
| test_loader: DataLoader, | |
| checkpoint_path: str | Path, | |
| output_dir: str | Path, | |
| device: torch.device, | |
| class_names: Sequence[str] = DEFAULT_CLASS_NAMES, | |
| run_metadata: dict[str, Any] | None = None, | |
| ) -> dict[str, Any]: | |
| """Evaluate a frozen validation-selected checkpoint exactly once.""" | |
| checkpoint = Path(checkpoint_path).resolve() | |
| if not checkpoint.is_file(): | |
| raise FileNotFoundError(f"Checkpoint does not exist: {checkpoint}") | |
| model.load_state_dict(load_file(str(checkpoint), device=str(device)), strict=True) | |
| result = predict_model(model, test_loader, device) | |
| metrics = compute_classification_metrics( | |
| result.true_labels, result.predicted_labels, class_names=class_names | |
| ) | |
| metrics.update( | |
| { | |
| "evaluated_split": "test", | |
| "checkpoint": str(checkpoint), | |
| "checkpoint_selection": "minimum_validation_loss", | |
| "parameter_count": sum(parameter.numel() for parameter in model.parameters()), | |
| "runtime": {"inference_seconds": result.inference_seconds}, | |
| "environment": package_environment(device), | |
| } | |
| ) | |
| if run_metadata: | |
| metrics["run"] = run_metadata | |
| destination = Path(output_dir).resolve() | |
| destination.mkdir(parents=True, exist_ok=True) | |
| write_predictions_csv( | |
| destination / "predictions.csv", | |
| result.filenames, | |
| result.true_labels, | |
| result.predicted_labels, | |
| result.probabilities, | |
| ) | |
| return write_json(destination / "metrics.json", metrics) | |
| def passes_publication_gate( | |
| good_test_macro_f1: Decimal, | |
| bad_test_macro_f1: Decimal, | |
| ) -> bool: | |
| """Apply the exact absolute 0.10 frozen-test macro-F1 threshold.""" | |
| return good_test_macro_f1 - bad_test_macro_f1 >= PUBLICATION_THRESHOLD | |
| def write_comparison_json( | |
| path: str | Path, | |
| good_test_macro_f1: Decimal, | |
| bad_test_macro_f1: Decimal, | |
| ) -> dict[str, Any]: | |
| """Write the exact Decimal comparison used for the publication decision.""" | |
| good = Decimal(good_test_macro_f1) | |
| bad = Decimal(bad_test_macro_f1) | |
| payload = { | |
| "fold": 0, | |
| "metric": "test_macro_f1", | |
| "good": str(good), | |
| "bad": str(bad), | |
| "gap": str(good - bad), | |
| "publication_threshold": str(PUBLICATION_THRESHOLD), | |
| "passes_publication_gate": passes_publication_gate(good, bad), | |
| } | |
| return write_json(path, payload) | |
| def resolve_device(requested: str = "auto") -> torch.device: | |
| if requested == "auto": | |
| return torch.device("cuda" if torch.cuda.is_available() else "cpu") | |
| device = torch.device(requested) | |
| if device.type == "cuda" and not torch.cuda.is_available(): | |
| raise RuntimeError("CUDA was requested but is not available") | |
| return device | |
| def run_test_evaluation( | |
| run_kind: str, | |
| labels_dir: str | Path, | |
| images_dir: str | Path, | |
| checkpoint_path: str | Path, | |
| output_dir: str | Path, | |
| device: str = "auto", | |
| num_workers: int = 4, | |
| ) -> dict[str, Any]: | |
| """Build only the official test loader and evaluate one frozen run.""" | |
| if run_kind not in {"good", "bad"}: | |
| raise ValueError("run_kind must be 'good' or 'bad'") | |
| labels_path = Path(labels_dir).resolve() | |
| images_path = Path(images_dir).resolve() | |
| output_path = Path(output_dir).resolve() | |
| selected_device = resolve_device(device) | |
| splits = load_fold_splits(labels_dir=labels_path, images_dir=images_path, fold=0) | |
| transforms = build_comparison_transforms() | |
| test_loader = make_loader( | |
| splits["test"], | |
| transforms[f"{run_kind}_test"], | |
| shuffle=False, | |
| seed=2026, | |
| batch_size=64, | |
| num_workers=num_workers, | |
| pin_memory=selected_device.type == "cuda", | |
| ) | |
| model = SmallDeepWeedsCNN(num_classes=9, in_channels=3 if run_kind == "good" else 1) | |
| return evaluate_checkpoint( | |
| model, | |
| test_loader, | |
| checkpoint_path, | |
| output_path, | |
| selected_device, | |
| run_metadata={ | |
| "kind": run_kind, | |
| "seed": 2026, | |
| "fold": 0, | |
| "test_fold_csv": str((labels_path / "test_subset0.csv").resolve()), | |
| }, | |
| ) | |
| def _parser() -> argparse.ArgumentParser: | |
| parser = argparse.ArgumentParser(description=__doc__) | |
| subparsers = parser.add_subparsers(dest="command", required=True) | |
| evaluate = subparsers.add_parser("evaluate", help="evaluate one frozen checkpoint") | |
| evaluate.add_argument("--run-kind", choices=("good", "bad"), required=True) | |
| evaluate.add_argument("--labels-dir", type=Path, required=True) | |
| evaluate.add_argument("--images-dir", type=Path, required=True) | |
| evaluate.add_argument("--checkpoint", type=Path, required=True) | |
| evaluate.add_argument("--output-dir", type=Path, required=True) | |
| evaluate.add_argument("--device", default="auto") | |
| evaluate.add_argument("--num-workers", type=int, default=4) | |
| compare = subparsers.add_parser("compare", help="write the frozen test comparison") | |
| compare.add_argument("--good-metrics", type=Path, required=True) | |
| compare.add_argument("--bad-metrics", type=Path, required=True) | |
| compare.add_argument("--output", type=Path, required=True) | |
| return parser | |
| def main(argv: Sequence[str] | None = None) -> int: | |
| args = _parser().parse_args(argv) | |
| if args.command == "evaluate": | |
| metrics = run_test_evaluation( | |
| run_kind=args.run_kind, | |
| labels_dir=args.labels_dir, | |
| images_dir=args.images_dir, | |
| checkpoint_path=args.checkpoint, | |
| output_dir=args.output_dir, | |
| device=args.device, | |
| num_workers=args.num_workers, | |
| ) | |
| print(json.dumps(metrics, indent=2, allow_nan=False)) | |
| return 0 | |
| good_metrics = json.loads(args.good_metrics.resolve().read_text(encoding="utf-8")) | |
| bad_metrics = json.loads(args.bad_metrics.resolve().read_text(encoding="utf-8")) | |
| comparison = write_comparison_json( | |
| args.output, | |
| Decimal(str(good_metrics["macro_f1"])), | |
| Decimal(str(bad_metrics["macro_f1"])), | |
| ) | |
| print(json.dumps(comparison, indent=2, allow_nan=False)) | |
| return 0 if comparison["passes_publication_gate"] else 2 | |
| if __name__ == "__main__": | |
| sys.exit(main()) | |