from __future__ import annotations from pathlib import Path from typing import Any import numpy as np import xarray as xr from scipy.ndimage import binary_dilation from .input_data import ConcatInput, L2AIIInput from .utils import future_offsets, inv_zscore, load_stats, open_memmap, parse_time, zscore class CILabel: source_name = "ci" def __init__(self, config: dict[str, Any]): self.config = dict(config) self.root = Path(self.config["root"]) self.expected_hw = self.config.get("shape_hw") if self.expected_hw is not None: self.expected_hw = tuple(int(v) for v in self.expected_hw) self._xr_open_kwargs = dict(decode_cf=False, mask_and_scale=False, decode_times=False) @property def row_shape(self) -> tuple[int, int]: if not self.expected_hw: raise ValueError("ci.shape_hw must be configured or inferred before memmap use") return tuple(self.expected_hw) def path_for_time(self, timestamp: str): ts = parse_time(timestamp).strftime("%Y%m%d%H%M") day = ts[:8] nc = self.root / day / f"{ts}_label.nc" if nc.exists(): return nc npz = self.root / day / f"{ts}_layers.npz" if npz.exists(): return npz return None def load_label(self, timestamp: str) -> np.ndarray: path = self.path_for_time(timestamp) if path is None: raise FileNotFoundError(f"CI label not found for {timestamp}") if path.suffix == ".npz": with np.load(path, allow_pickle=False) as data: if "masks" not in data: raise KeyError(f"'masks' key not found in {path}") masks = data["masks"] if masks.ndim != 3: raise ValueError(f"npz masks must be (N,H,W), got {masks.shape}: {path}") label = (masks > 0).any(axis=0).astype(np.uint8) else: with xr.open_dataset(path, **self._xr_open_kwargs) as ds: var = list(ds.data_vars)[0] arr = ds[var].values if arr.ndim != 2: raise ValueError(f"CI label must be 2D, got {arr.shape}: {path}") label = (~np.isnan(arr)).astype(np.uint8) if self.expected_hw is None: self.expected_hw = tuple(int(v) for v in label.shape) if tuple(label.shape) != tuple(self.expected_hw): raise ValueError(f"CI shape mismatch: {label.shape} != {self.expected_hw}") return label @staticmethod def smooth(label: np.ndarray, base: float = 0.5, radius: int = 3) -> np.ndarray: base = float(base) radius = int(radius) if not (0.0 < base < 1.0): raise ValueError(f"base must be between 0 and 1, got {base}") if radius < 0: raise ValueError(f"radius must be >= 0, got {radius}") positive = np.asarray(label) > 0.5 smoothed = np.zeros(positive.shape, dtype=np.float32) smoothed[positive] = 1.0 if radius == 0 or not positive.any(): return smoothed structure = np.ones((3, 3), dtype=bool) prev = positive for distance in range(1, radius + 1): dilated = binary_dilation(positive, structure=structure, iterations=distance) ring = dilated & ~prev smoothed[ring] = base ** distance prev = dilated return smoothed def open_memmap(self, dat_path: str | Path, n_rows: int, dtype: str = "uint8", mode: str = "r") -> np.memmap: return open_memmap(dat_path, dtype, (int(n_rows), *self.row_shape), mode=mode) def load_memmap_row(self, dat_path: str | Path, row_idx: int, n_rows: int, dtype: str = "uint8") -> np.ndarray: mm = self.open_memmap(dat_path, n_rows=n_rows, dtype=dtype, mode="r") return np.asarray(mm[int(row_idx)]) class BTLabel: source_name = "bt" file_dtype = "float16" DEFAULT_CONDITIONS = { "KI": {"op": ">=", "value": 30.0}, "LI": {"op": "<=", "value": -2.0}, "SI": {"op": "<=", "value": 2.0}, "CAPE": {"op": ">=", "value": 500.0}, "TTI": {"op": ">=", "value": 42.0}, } def __init__(self, config: dict[str, Any], concat_input: ConcatInput, l2_input: L2AIIInput): self.config = dict(config) self.concat_input = concat_input self.l2_input = l2_input self.future_minutes = int(self.config.get("future_minutes", 60)) self.interval_minutes = int(self.config.get("interval_minutes", 10)) self.lead_minutes = list(self.config.get("lead_minutes") or future_offsets(self.future_minutes, self.interval_minutes)) self.var_name = str(self.config.get("var_name", "ir105")) self.mask_var_name = str(self.config.get("mask_var_name", self.var_name)) self.bt_threshold_k = float(self.config.get("bt_threshold_k", 233.0)) self.expansion_km = float(self.config.get("expansion_km", 50.0)) self.pixel_size_km = float(self.config.get("pixel_size_km", 2.0)) self.conditions = dict(self.DEFAULT_CONDITIONS) self.conditions.update(self.config.get("aii_conditions", {})) self.stats_path = self.config.get("stats_path") or self.concat_input.stats_path self.stats = load_stats(self.stats_path) self.normalization = str(self.config.get("normalization", "zscore")).lower() self.eps = float(self.config.get("eps", 1e-6)) self.expected_hw = self.config.get("shape_hw") or self.concat_input.expected_hw if self.expected_hw is not None: self.expected_hw = tuple(int(v) for v in self.expected_hw) @property def row_shape(self) -> tuple[int, int, int]: if not self.expected_hw: raise ValueError("BT shape_hw must be configured or inferred before memmap use") h, w = self.expected_hw return (1, h, w) @property def mask_row_shape(self) -> tuple[int, int]: if not self.expected_hw: raise ValueError("BT shape_hw must be configured or inferred before memmap use") return tuple(self.expected_hw) def _condition_mask(self, arr: np.ndarray, op: str, value: float) -> np.ndarray: arr = np.asarray(arr, dtype=np.float32) finite = np.isfinite(arr) if op == ">=": return finite & (arr >= float(value)) if op == ">": return finite & (arr > float(value)) if op == "<=": return finite & (arr <= float(value)) if op == "<": return finite & (arr < float(value)) raise ValueError(f"unsupported condition op: {op}") def circular_footprint(self) -> np.ndarray: radius_pixels = int(np.ceil(self.expansion_km / self.pixel_size_km)) yy, xx = np.ogrid[-radius_pixels : radius_pixels + 1, -radius_pixels : radius_pixels + 1] distance_km = np.sqrt(xx**2 + yy**2) * self.pixel_size_km return distance_km <= self.expansion_km def build_mask(self, timestamp: str, shape_hw: tuple[int, int]) -> np.ndarray: aii = self.l2_input.load_raw_dict(timestamp) masks = [] for name, rule in self.conditions.items(): if name not in aii: raise KeyError(f"BT mask requires L2 variable {name!r}") arr = np.asarray(aii[name], dtype=np.float32) if tuple(arr.shape) != tuple(shape_hw): raise ValueError(f"AII/BT shape mismatch for {name}: {arr.shape} != {shape_hw}") masks.append(self._condition_mask(arr, str(rule["op"]), float(rule["value"]))) aii_mask = np.logical_and.reduce(masks) start_data = self.concat_input.load_raw_dict(timestamp) start_bt = self.concat_input.get_variable(start_data, self.mask_var_name) if tuple(start_bt.shape) != tuple(shape_hw): raise ValueError(f"start mask-variable shape mismatch: {start_bt.shape} != {shape_hw}") cold = np.isfinite(start_bt) & (start_bt <= self.bt_threshold_k) expanded = binary_dilation(cold, structure=self.circular_footprint()) if self.expansion_km > 0 else cold return aii_mask & ~expanded def normalize_bt(self, arr: np.ndarray) -> np.ndarray: arr = np.asarray(arr, dtype=np.float32) if self.normalization in {"none", "raw", "false"}: return arr if self.normalization == "zscore": stat = self.stats[self.var_name] return zscore(arr, stat["mean"], stat["std"], self.eps) raise ValueError(f"unsupported BT normalization: {self.normalization}") def denormalize(self, arr: np.ndarray) -> np.ndarray: arr = np.asarray(arr, dtype=np.float32) if self.normalization in {"none", "raw", "false"}: return arr stat = self.stats[self.var_name] return inv_zscore(arr, stat["mean"], stat["std"], self.eps) def load_label(self, timestamp: str) -> np.ndarray: data = self.concat_input.load_raw_dict(timestamp) bt = self.concat_input.get_variable(data, self.var_name) if self.expected_hw is None: self.expected_hw = tuple(int(v) for v in bt.shape) if tuple(bt.shape) != tuple(self.expected_hw): raise ValueError(f"BT shape mismatch: {timestamp} {bt.shape} != {self.expected_hw}") normalized = self.normalize_bt(bt) return normalized[np.newaxis, ...].astype(np.float32, copy=False) def load_mask(self, timestamp: str) -> np.ndarray: data = self.concat_input.load_raw_dict(timestamp) bt = self.concat_input.get_variable(data, self.var_name) if self.expected_hw is None: self.expected_hw = tuple(int(v) for v in bt.shape) if tuple(bt.shape) != tuple(self.expected_hw): raise ValueError(f"BT mask shape mismatch: {timestamp} {bt.shape} != {self.expected_hw}") return self.build_mask(timestamp, tuple(self.expected_hw)).astype(np.uint8, copy=False) def open_memmap(self, dat_path: str | Path, n_rows: int, dtype: str | None = None, mode: str = "r") -> np.memmap: return open_memmap(dat_path, dtype or self.file_dtype, (int(n_rows), *self.row_shape), mode=mode) def load_memmap_row(self, dat_path: str | Path, row_idx: int, n_rows: int, dtype: str | None = None) -> np.ndarray: mm = self.open_memmap(dat_path, n_rows=n_rows, dtype=dtype or self.file_dtype, mode="r") return np.asarray(mm[int(row_idx)], dtype=np.float32)