ci-net / code /validation /src /validator.py
lsh9034's picture
Add files using upload-large-folder tool
7da2ecb verified
Raw History Blame Contribute Delete
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
@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,
}