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