File size: 5,818 Bytes
7da2ecb | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 | """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)))
|