"""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)) @dataclass(frozen=True) 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())