Download code/validation/src/validator.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 15.9 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/validation/src/validator.py
- Command line
-
hf download hf://lsh9034/ci-net/code/validation/src/validator.py
-
curl -L -o validator.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/validation/src/validator.py
15.9 kB
| """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 | |
| 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") | |
| 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 | |
| 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, | |
| } | |