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