"""Threshold-sweep validation helpers.""" from __future__ import annotations from collections import defaultdict from dataclasses import dataclass from datetime import datetime from typing import Any import numpy as np from tqdm import tqdm from .clusterers import Clusterer from .loaders import CloudLabelLoader, CloudTarget, PredictionProvider from .utils import format_dt from .validator import RawValidationResult, Validator @dataclass class ThresholdSweepResult: raw_results: dict[float, RawValidationResult] class ThresholdSweepValidator: """Validate one target set for several model thresholds with one I/O pass.""" def __init__( self, targets: list[CloudTarget], label_loader: CloudLabelLoader, prediction_provider: PredictionProvider, clusterers: dict[float, Clusterer], pixel_size_km: float = 2.0, label_buffer_km: float = 0.0, model_false_buffer_km: float = 0.0, leadtime_min: int = 10, leadtime_max: int = 120, time_step: int = 10, show_progress: bool = True, buffer_backend: str = "auto", ): if not clusterers: raise ValueError("At least one threshold clusterer is required") self.targets = targets self.label_loader = label_loader self.prediction_provider = prediction_provider self.clusterers = {float(threshold): clusterer for threshold, clusterer in clusterers.items()} self.thresholds = sorted(self.clusterers) self.show_progress = bool(show_progress) # Reuse the established matching, leadtime, and summary logic exactly. self._helper = Validator( targets=targets, label_loader=label_loader, prediction_provider=prediction_provider, clusterer=self.clusterers[self.thresholds[0]], pixel_size_km=pixel_size_km, label_buffer_km=label_buffer_km, model_false_buffer_km=model_false_buffer_km, leadtime_min=leadtime_min, leadtime_max=leadtime_max, time_step=time_step, buffer_backend=buffer_backend, ) def evaluate_raw(self) -> ThresholdSweepResult: targets_by_dt: dict[datetime, list[CloudTarget]] = defaultdict(list) for target in self.targets: targets_by_dt[target.dt].append(target) label_records_by_threshold: dict[float, list[dict[str, Any]]] = { threshold: [] for threshold in self.thresholds } model_records_by_threshold: dict[float, list[dict[str, Any]]] = { threshold: [] for threshold in self.thresholds } missing_predictions: list[dict[str, Any]] = [] dts = sorted(targets_by_dt) iterator = tqdm(dts, desc="Threshold sweep timesteps", dynamic_ncols=True) if self.show_progress else dts for dt in iterator: dt_targets = targets_by_dt[dt] label_arr = self.label_loader.load(dt) try: field = self.prediction_provider.load(dt) except FileNotFoundError as exc: missing_predictions.append( { "time": format_dt(dt), "reason": "missing_prediction", "message": str(exc), "num_labels": len(dt_targets), "cloud_ids": [target.cloud_id for target in dt_targets], } ) for threshold in self.thresholds: for target in dt_targets: label_records_by_threshold[threshold].append( self._helper._base_label_record( target, status="impossible", matched_cluster_ids=[], prediction_path=None, label_pixel_count=None, reason="missing_prediction", ) ) continue if field.data.shape != label_arr.shape: raise ValueError( f"Shape mismatch at {format_dt(dt)}: prediction={field.data.shape}, label={label_arr.shape}" ) target_masks = self._helper._build_target_masks(label_arr, dt_targets) target_distance_maps = self._helper._build_target_distance_maps(target_masks) targets_by_cloud_id = {target.cloud_id: target for target in dt_targets} for threshold in self.thresholds: clusters = self.clusterers[threshold].cluster(field.data, field.valid_mask) for target in dt_targets: label_mask = target_masks[target.cloud_id] label_distance = target_distance_maps[target.cloud_id] matched_clusters = self._helper._matched_clusters_for_label( label_mask, label_distance, clusters, ) status = "hit" if matched_clusters else "miss" label_records_by_threshold[threshold].append( self._helper._base_label_record( target, status=status, matched_cluster_ids=[cluster.cluster_id for cluster in matched_clusters], prediction_path=field.path, label_pixel_count=int(label_mask.sum()), reason="matched" if matched_clusters else "no_matching_model_cluster", ) ) model_records_by_threshold[threshold].extend( self._helper._count_model_clusters( dt=dt, clusters=clusters, target_masks=target_masks, target_distance_maps=target_distance_maps, targets_by_cloud_id=targets_by_cloud_id, prediction_path=field.path, ) ) raw_results = { threshold: RawValidationResult( label_records=label_records_by_threshold[threshold], model_cluster_records=model_records_by_threshold[threshold], missing_predictions=list(missing_predictions), ) for threshold in self.thresholds } return ThresholdSweepResult(raw_results=raw_results) def apply_leadtime_mode(self, raw_label_records: list[dict[str, Any]], leadtime_mode: str) -> list[dict[str, Any]]: return self._helper.apply_leadtime_mode(raw_label_records, leadtime_mode) def summarize( self, label_records: list[dict[str, Any]], model_cluster_records: list[dict[str, Any]], ) -> dict[str, Any]: return self._helper.summarize(label_records, model_cluster_records)