"""JSON, label, and prediction loading.""" from __future__ import annotations import json from dataclasses import dataclass from datetime import datetime from pathlib import Path from typing import Any import numpy as np import xarray as xr from .utils import day_str, format_dt, parse_cloud_id @dataclass(frozen=True) class CloudTarget: mature_id: str mature_dt: datetime mature_number: int cloud_id: str dt: datetime number: int should_validate: bool leadtime: int def to_dict(self) -> dict[str, Any]: return { "mature_id": self.mature_id, "mature_time": format_dt(self.mature_dt), "mature_number": self.mature_number, "cloud_id": self.cloud_id, "cloud_time": format_dt(self.dt), "cloud_number": self.number, "should_validate": self.should_validate, "leadtime": self.leadtime, } @dataclass class PredictionField: data: np.ndarray valid_mask: np.ndarray path: str class ValidationJsonLoader: def __init__(self, json_path: str | Path): self.json_path = Path(json_path) def load_raw(self) -> dict[str, dict[str, bool]]: with open(self.json_path, "r", encoding="utf-8") as f: data = json.load(f) if not isinstance(data, dict): raise ValueError(f"Validation JSON must be a dict: {self.json_path}") return data def load_targets( self, target_filter: str, leadtime_min: int | None = None, leadtime_max: int | None = None, max_cases: int | None = None, ) -> list[CloudTarget]: if target_filter not in {"true_only", "all"}: raise ValueError(f"Unsupported target_filter: {target_filter}") data = self.load_raw() targets: list[CloudTarget] = [] seen_cloud_ids: set[str] = set() for mature_id, past_clouds in data.items(): if not isinstance(past_clouds, dict): continue mature_dt, mature_number = parse_cloud_id(mature_id) for cloud_id, should_validate in sorted(past_clouds.items(), key=lambda item: item[0]): should_validate = bool(should_validate) if target_filter == "true_only" and not should_validate: continue if cloud_id in seen_cloud_ids: continue cloud_dt, cloud_number = parse_cloud_id(cloud_id) leadtime = int((mature_dt - cloud_dt).total_seconds() // 60) if leadtime_min is not None and leadtime < leadtime_min: continue if leadtime_max is not None and leadtime > leadtime_max: continue targets.append( CloudTarget( mature_id=mature_id, mature_dt=mature_dt, mature_number=mature_number, cloud_id=cloud_id, dt=cloud_dt, number=cloud_number, should_validate=should_validate, leadtime=leadtime, ) ) seen_cloud_ids.add(cloud_id) if max_cases is not None and len(targets) >= max_cases: return targets return targets class CloudLabelLoader: def __init__(self, temporal_overlapping_dir: str | Path): self.temporal_overlapping_dir = Path(temporal_overlapping_dir) self._cache: dict[str, np.ndarray] = {} def label_path(self, dt: datetime) -> Path: dt_str = format_dt(dt) return self.temporal_overlapping_dir / day_str(dt) / f"{dt_str}_label.nc" def load(self, dt: datetime) -> np.ndarray: dt_str = format_dt(dt) if dt_str in self._cache: return self._cache[dt_str] path = self.label_path(dt) if not path.exists(): raise FileNotFoundError(f"Label file not found: {path}") with xr.open_dataset(path) as ds: label = ds["label"].values self._cache[dt_str] = label return label class PredictionProvider: name = "Base" def __init__(self, root_dir: str | Path): self.root_dir = Path(root_dir) def path_for_dt(self, dt: datetime) -> Path: raise NotImplementedError def load(self, dt: datetime) -> PredictionField: path = self.path_for_dt(dt) if not path.exists(): raise FileNotFoundError(f"Prediction file not found: {path}") data = np.load(path, allow_pickle=True).squeeze().astype(float) valid_mask = np.isfinite(data) return PredictionField(data=data, valid_mask=valid_mask, path=str(path)) class ModelProvider(PredictionProvider): name = "Model" def __init__(self, root_dir: str | Path, use_masked: bool = False): super().__init__(root_dir) self.use_masked = bool(use_masked) def path_for_dt(self, dt: datetime) -> Path: dt_str = format_dt(dt) suffix = "_masked" if self.use_masked else "" return self.root_dir / day_str(dt) / f"pred_{dt_str}{suffix}.npy" def create_prediction_provider(config: dict[str, Any]) -> PredictionProvider: source = config.get("data_source", "Model") if source != "Model": raise ValueError("The public release validates CI-Net model outputs only") model_dirs = config["paths"]["model_output_dirs"] if source not in model_dirs: raise ValueError(f"Missing model output dir for data_source={source}") model_config = (config.get("providers") or {}).get("Model", {}) return ModelProvider(model_dirs[source], use_masked=bool(model_config.get("use_masked", False)))