"""Object-based validation loop.""" 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 scipy.ndimage import distance_transform_edt from tqdm import tqdm from .clusterers import Clusterer, ModelCluster from .loaders import CloudLabelLoader, CloudTarget, PredictionProvider from .metrics import compute_scores from .utils import dilate, format_dt, km_to_pixels, use_edt_backend @dataclass class RawValidationResult: label_records: list[dict[str, Any]] model_cluster_records: list[dict[str, Any]] missing_predictions: list[dict[str, Any]] class Validator: """Validate backtracked cloud labels against clustered model predictions.""" def __init__( self, targets: list[CloudTarget], label_loader: CloudLabelLoader, prediction_provider: PredictionProvider, clusterer: 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, buffer_backend: str = "auto", ): self.targets = targets self.label_loader = label_loader self.prediction_provider = prediction_provider self.clusterer = clusterer self.pixel_size_km = float(pixel_size_km) self.label_buffer_pixels = km_to_pixels(label_buffer_km, self.pixel_size_km) self.model_false_buffer_pixels = km_to_pixels(model_false_buffer_km, self.pixel_size_km) self.leadtime_min = int(leadtime_min) self.leadtime_max = int(leadtime_max) self.time_step = int(time_step) self.buffer_backend = str(buffer_backend or "auto") @property def leadtimes(self) -> list[int]: return list(range(self.leadtime_max, self.leadtime_min - 1, -self.time_step)) def evaluate_raw(self) -> RawValidationResult: targets_by_dt: dict[datetime, list[CloudTarget]] = defaultdict(list) for target in self.targets: targets_by_dt[target.dt].append(target) label_records: list[dict[str, Any]] = [] model_cluster_records: list[dict[str, Any]] = [] missing_predictions: list[dict[str, Any]] = [] for dt in tqdm(sorted(targets_by_dt), desc="Validate timesteps", dynamic_ncols=True): 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 target in dt_targets: label_records.append( self._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}" ) clusters = self.clusterer.cluster(field.data, field.valid_mask) target_masks = self._build_target_masks(label_arr, dt_targets) target_distance_maps = self._build_target_distance_maps(target_masks) for target in dt_targets: label_mask = target_masks[target.cloud_id] label_distance = target_distance_maps[target.cloud_id] matched_clusters = self._matched_clusters_for_label(label_mask, label_distance, clusters) status = "hit" if matched_clusters else "miss" label_records.append( self._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_cluster_records.extend( self._count_model_clusters( dt=dt, clusters=clusters, target_masks=target_masks, target_distance_maps=target_distance_maps, targets_by_cloud_id={target.cloud_id: target for target in dt_targets}, prediction_path=field.path, ) ) return RawValidationResult( label_records=label_records, model_cluster_records=model_cluster_records, missing_predictions=missing_predictions, ) def _build_target_masks(self, label_arr: np.ndarray, targets: list[CloudTarget]) -> dict[str, np.ndarray]: masks: dict[str, np.ndarray] = {} for target in targets: mask = label_arr == target.number if not np.any(mask): raise ValueError(f"Backtracked cloud {target.cloud_id} not found in label") masks[target.cloud_id] = mask return masks def _build_target_distance_maps(self, target_masks: dict[str, np.ndarray]) -> dict[str, np.ndarray]: return { cloud_id: distance_transform_edt(~mask.astype(bool, copy=False)) for cloud_id, mask in target_masks.items() if np.any(mask) } def _matches_distance(self, cluster_mask: np.ndarray, target_distance: np.ndarray, radius_pixels: int) -> bool: if not np.any(cluster_mask): return False return bool(float(np.min(target_distance[cluster_mask.astype(bool, copy=False)])) <= int(radius_pixels)) def _matched_clusters_for_label( self, label_mask: np.ndarray, label_distance: np.ndarray, clusters: list[ModelCluster], ) -> list[ModelCluster]: if use_edt_backend(self.label_buffer_pixels, self.buffer_backend): return [ cluster for cluster in clusters if self._matches_distance(cluster.mask, label_distance, self.label_buffer_pixels) ] label_match_mask = dilate(label_mask, self.label_buffer_pixels) return [cluster for cluster in clusters if self._touches(label_match_mask, cluster.mask)] def _matched_cloud_ids_for_cluster( self, cluster_mask: np.ndarray, target_masks: dict[str, np.ndarray], target_distance_maps: dict[str, np.ndarray], ) -> list[str]: if use_edt_backend(self.model_false_buffer_pixels, self.buffer_backend): return sorted( cloud_id for cloud_id, distances in target_distance_maps.items() if self._matches_distance(cluster_mask, distances, self.model_false_buffer_pixels) ) cluster_match_mask = dilate(cluster_mask, self.model_false_buffer_pixels) return sorted( cloud_id for cloud_id, mask in target_masks.items() if self._touches(cluster_match_mask, mask) ) def _base_label_record( self, target: CloudTarget, status: str, matched_cluster_ids: list[int], prediction_path: str | None, label_pixel_count: int | None, reason: str, ) -> dict[str, Any]: record = target.to_dict() record.update( { "raw_status": status, "status": status, "reason": reason, "matched_cluster_ids": matched_cluster_ids, "matched_cluster_count": len(matched_cluster_ids), "prediction_path": prediction_path, "label_pixel_count": label_pixel_count, "label_buffer_pixels": self.label_buffer_pixels, "label_buffer_km": self.label_buffer_pixels * self.pixel_size_km, } ) return record def _count_model_clusters( self, dt: datetime, clusters: list[ModelCluster], target_masks: dict[str, np.ndarray], target_distance_maps: dict[str, np.ndarray] | None, targets_by_cloud_id: dict[str, CloudTarget], prediction_path: str, ) -> list[dict[str, Any]]: records: list[dict[str, Any]] = [] if target_distance_maps is None: target_distance_maps = self._build_target_distance_maps(target_masks) for cluster in clusters: matched_cloud_ids = self._matched_cloud_ids_for_cluster( cluster.mask, target_masks, target_distance_maps, ) is_false = len(matched_cloud_ids) == 0 nearest_label = ( self._nearest_target_label(cluster.mask, target_distance_maps, targets_by_cloud_id) if is_false else None ) records.append( { "time": format_dt(dt), "cluster_id": int(cluster.cluster_id), "cluster_key": f"{format_dt(dt)}_{cluster.cluster_id}", "status": "false" if is_false else "model_hit", "is_false": is_false, "matched_cloud_ids": matched_cloud_ids, "matched_cloud_count": len(matched_cloud_ids), "pixel_count": int(cluster.pixel_count), "raw_pixel_count": int(cluster.raw_pixel_count), "model_false_buffer_pixels": self.model_false_buffer_pixels, "model_false_buffer_km": self.model_false_buffer_pixels * self.pixel_size_km, "assigned_false_cloud_id": nearest_label["cloud_id"] if nearest_label else None, "assigned_false_leadtime": nearest_label["leadtime"] if nearest_label else None, "assigned_false_distance_pixels": nearest_label["distance_pixels"] if nearest_label else None, "assigned_false_distance_km": nearest_label["distance_km"] if nearest_label else None, "prediction_path": prediction_path, } ) return records def _nearest_target_label( self, cluster_mask: np.ndarray, target_distance_maps: dict[str, np.ndarray], targets_by_cloud_id: dict[str, CloudTarget], ) -> dict[str, Any] | None: if not np.any(cluster_mask) or not target_distance_maps: return None best: dict[str, Any] | None = None cluster_mask = cluster_mask.astype(bool, copy=False) for cloud_id in sorted(target_distance_maps): distances = target_distance_maps[cloud_id] distance_pixels = float(np.min(distances[cluster_mask])) target = targets_by_cloud_id[cloud_id] candidate = { "cloud_id": cloud_id, "leadtime": int(target.leadtime), "distance_pixels": distance_pixels, "distance_km": distance_pixels * self.pixel_size_km, } if best is None or distance_pixels < best["distance_pixels"]: best = candidate return best @staticmethod def _touches(mask_a: np.ndarray, mask_b: np.ndarray) -> bool: return bool(np.any(mask_a & mask_b)) def apply_leadtime_mode(self, raw_label_records: list[dict[str, Any]], leadtime_mode: str) -> list[dict[str, Any]]: if leadtime_mode == "exact": return [dict(record, status=record["raw_status"], leadtime_mode="exact") for record in raw_label_records] if leadtime_mode != "accumulate": raise ValueError(f"Unsupported leadtime_mode: {leadtime_mode}") records = [dict(record, leadtime_mode="accumulate") for record in raw_label_records] by_mature: dict[str, list[dict[str, Any]]] = defaultdict(list) for record in records: by_mature[record["mature_id"]].append(record) for mature_records in by_mature.values(): possible = [record for record in mature_records if record["raw_status"] != "impossible"] hit_leadtimes = [int(record["leadtime"]) for record in possible if record["raw_status"] == "hit"] first_hit_leadtime = max(hit_leadtimes) if hit_leadtimes else None for record in mature_records: if record["raw_status"] == "impossible": record["status"] = "impossible" record["accumulate_first_hit_leadtime"] = first_hit_leadtime elif first_hit_leadtime is not None and int(record["leadtime"]) <= first_hit_leadtime: record["status"] = "hit" record["reason"] = "accumulated_from_first_hit" record["accumulate_first_hit_leadtime"] = first_hit_leadtime else: record["status"] = "miss" record["accumulate_first_hit_leadtime"] = first_hit_leadtime return records def summarize( self, label_records: list[dict[str, Any]], model_cluster_records: list[dict[str, Any]], ) -> dict[str, Any]: hits = sum(1 for record in label_records if record["status"] == "hit") misses = sum(1 for record in label_records if record["status"] == "miss") impossible = sum(1 for record in label_records if record["status"] == "impossible") falses = sum(1 for record in model_cluster_records if record["is_false"]) model_hits = sum(1 for record in model_cluster_records if not record["is_false"]) leadtime_metrics: dict[int, dict[str, Any]] = {} for leadtime in self.leadtimes: lt_records = [record for record in label_records if int(record["leadtime"]) == leadtime] lt_hits = sum(1 for record in lt_records if record["status"] == "hit") lt_misses = sum(1 for record in lt_records if record["status"] == "miss") lt_impossible = sum(1 for record in lt_records if record["status"] == "impossible") lt_falses = sum( 1 for record in model_cluster_records if record["is_false"] and record.get("assigned_false_leadtime") == leadtime ) lt_scores = compute_scores(lt_hits, lt_misses, lt_falses) lt_scores.update( { "impossible": int(lt_impossible), "valid_labels": int(lt_hits + lt_misses), } ) leadtime_metrics[int(leadtime)] = lt_scores scores = compute_scores(hits, misses, falses) scores.update( { "valid_labels": int(hits + misses), "impossible": int(impossible), "total_input_labels": int(len(label_records)), "model_hit_clusters": int(model_hits), "false_model_clusters": int(falses), "total_model_clusters": int(len(model_cluster_records)), } ) return { "total": scores, "leadtime": leadtime_metrics, }