from __future__ import annotations from pathlib import Path from typing import Any import math import re import numpy as np import xarray as xr from .utils import ( as_path_list, format_time, inv_zscore, load_stats, normalize_mode, open_memmap, parse_time, safe_nanmean, zscore, ) class BaseInput: source_name = "base" file_dtype = "float16" def __init__(self, config: dict[str, Any]): self.config = dict(config) self.roots = as_path_list(self.config["root"]) self.selected_vars = list(self.config.get("selected_vars", [])) self.stats_path = self.config.get("stats_path") self.stats = load_stats(self.stats_path) self.stats_name_map = { str(name): str(stats_name) for name, stats_name in self.config.get("stats_name_map", {}).items() } self.normalization = normalize_mode(self.config) self.eps = float(self.config.get("eps", 1e-6)) 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) self.invalid_fill = self._build_invalid_fill_config(self.config.get("invalid_fill")) self.transforms = self._build_transforms_config(self.config.get("transforms")) @property def channels(self) -> list[str]: return list(self.selected_vars) @property def row_shape(self) -> tuple[int, int, int]: if not self.expected_hw: raise ValueError(f"{self.source_name}.shape_hw must be configured or inferred before memmap use") h, w = self.expected_hw return (len(self.channels), h, w) def path_for_time(self, timestamp: str): raise NotImplementedError def load_raw_dict(self, timestamp: str) -> dict[str, np.ndarray]: raise NotImplementedError def _build_invalid_fill_config(self, config: Any) -> dict[str, Any]: if config is None: config = {"default": {"method": "keep"}} if isinstance(config, str): config = {"default": {"method": config}} default = config.get("default", {"method": "keep"}) if isinstance(default, str): default = {"method": default} variables = {} for name, rule in config.get("variables", {}).items(): variables[name] = {"method": rule} if isinstance(rule, str) else dict(rule) return {"default": dict(default), "variables": variables} def _build_transforms_config(self, config: Any) -> dict[str, list[dict[str, Any]]]: if config is None: return {"before_normalize": [], "after_normalize": []} if isinstance(config, list): config = {"before_normalize": config} before = config.get("before_normalize", config.get("pre_normalize", [])) after = config.get("after_normalize", config.get("post_normalize", [])) return { "before_normalize": [dict(item) for item in before], "after_normalize": [dict(item) for item in after], } def _invalid_fill_rule(self, name: str) -> dict[str, Any]: return self.invalid_fill["variables"].get(name, self.invalid_fill["default"]) def _fill_timing(self, name: str) -> str: rule = self._invalid_fill_rule(name) return str(rule.get("timing", "before_normalize")).lower() def should_normalize(self, name: str) -> bool: return True def should_fill_invalid(self, name: str) -> bool: return True def _fill_by_rule(self, name: str, arr: np.ndarray) -> np.ndarray: arr = np.asarray(arr, dtype=np.float32) mask = ~np.isfinite(arr) if not mask.any(): return arr rule = self._invalid_fill_rule(name) method = str(rule.get("method", "keep")).lower() if method == "keep": return arr if method == "zero": return np.where(mask, 0.0, arr).astype(np.float32, copy=False) if method == "constant": value = float(rule.get("value", 0.0)) return np.where(mask, value, arr).astype(np.float32, copy=False) if method == "scene_mean": sample = arr stride = rule.get("stride") if stride is not None and int(stride) >= 2: sample = arr[:: int(stride), :: int(stride)] return np.where(mask, safe_nanmean(sample), arr).astype(np.float32, copy=False) raise ValueError(f"unsupported invalid fill method for {self.source_name}:{name}: {method}") def fill_invalid(self, name: str, arr: np.ndarray) -> np.ndarray: return self._fill_by_rule(name, arr) def _apply_transform(self, name: str, arr: np.ndarray, transform: dict[str, Any]) -> np.ndarray: arr = np.asarray(arr, dtype=np.float32) transform_type = str(transform.get("type", transform.get("name", ""))).lower() if transform_type == "clip": min_value = transform.get("min", transform.get("lower", None)) max_value = transform.get("max", transform.get("upper", None)) return np.clip( arr, -np.inf if min_value is None else float(min_value), np.inf if max_value is None else float(max_value), ).astype(np.float32, copy=False) if transform_type == "clamp": min_value = transform.get("min", transform.get("lower", None)) max_value = transform.get("max", transform.get("upper", None)) return np.clip( arr, -np.inf if min_value is None else float(min_value), np.inf if max_value is None else float(max_value), ).astype(np.float32, copy=False) if transform_type in {"divide", "div"}: value = float(transform["value"]) if value == 0.0: raise ValueError(f"{self.source_name}:{name} divide transform value must be non-zero") return (arr / value).astype(np.float32, copy=False) if transform_type in {"multiply", "mul", "scale"}: return (arr * float(transform["value"])).astype(np.float32, copy=False) if transform_type in {"replace_value", "replace"}: old_value = float(transform.get("old", transform.get("from"))) new_value = float(transform.get("new", transform.get("to", 0.0))) atol = float(transform.get("atol", 0.0)) mask = np.isclose(arr, old_value, rtol=0.0, atol=atol) return np.where(mask, new_value, arr).astype(np.float32, copy=False) if transform_type == "log1p": if np.any(np.isfinite(arr) & (arr <= -1.0)): min_value = float(np.nanmin(arr)) raise ValueError( f"{self.source_name}:{name} log1p requires values greater than -1, got min={min_value}" ) return np.log1p(arr).astype(np.float32, copy=False) raise ValueError(f"unsupported transform for {self.source_name}:{name}: {transform_type!r}") def apply_transforms(self, name: str, arr: np.ndarray, timing: str) -> np.ndarray: out = np.asarray(arr, dtype=np.float32) for transform in self.transforms.get(timing, []): out = self._apply_transform(name, out, transform) return out def _normalize_one(self, name: str, arr: np.ndarray) -> np.ndarray: if self.normalization in {"none", "raw", "false"}: return arr.astype(np.float32, copy=False) if self.normalization == "zscore": stats_name = self.stats_name_map.get(name, name) if stats_name not in self.stats: raise KeyError( f"statistics for {stats_name!r} (input variable {name!r}) " f"not found in {self.stats_path}" ) stat = self.stats[stats_name] return zscore(arr.astype(np.float32, copy=False), stat["mean"], stat["std"], self.eps) raise ValueError(f"unsupported normalization mode: {self.normalization}") def denormalize(self, name: str, arr: np.ndarray) -> np.ndarray: if self.normalization in {"none", "raw", "false"}: return np.asarray(arr, dtype=np.float32) if self.normalization == "zscore": stats_name = self.stats_name_map.get(name, name) stat = self.stats[stats_name] return inv_zscore(np.asarray(arr, dtype=np.float32), stat["mean"], stat["std"], self.eps) raise ValueError(f"unsupported normalization mode: {self.normalization}") def load_frame(self, timestamp: str, normalize: bool = True) -> np.ndarray: data = self.load_raw_dict(timestamp) arrays = [] if not self.channels: raise ValueError( f"{self.source_name} has no channels configured. " "Set selected_vars for simple inputs, or raw_vars/physics_formulas " "for concat inputs. If this source is only used as a raw helper " "for another label, call load_raw_dict() instead of load_frame()." ) for name in self.channels: arr = self.get_variable(data, name) if arr.ndim != 2: raise ValueError(f"{self.source_name}:{name} must be 2D, got {arr.shape}") if self.expected_hw is None: self.expected_hw = tuple(int(v) for v in arr.shape) if tuple(arr.shape) != tuple(self.expected_hw): raise ValueError(f"shape mismatch for {self.source_name}:{name}: {arr.shape} != {self.expected_hw}") fill_timing = self._fill_timing(name) if self.should_fill_invalid(name) and fill_timing != "after_normalize": arr = self.fill_invalid(name, arr) arr = self.apply_transforms(name, arr, "before_normalize") if normalize and self.should_normalize(name): arr = self._normalize_one(name, arr) if self.should_fill_invalid(name) and fill_timing == "after_normalize": arr = self.fill_invalid(name, arr) arr = self.apply_transforms(name, arr, "after_normalize") arrays.append(arr.astype(np.float32, copy=False)) return np.stack(arrays, axis=0) def get_variable(self, data: dict[str, np.ndarray], name: str) -> np.ndarray: if name not in data: raise KeyError(f"{name!r} not found. available={sorted(data)}") return np.asarray(data[name], dtype=np.float32) 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, mode="r") return np.asarray(mm[int(row_idx)], dtype=np.float32) class ConcatInput(BaseInput): source_name = "concat" file_dtype = "float16" NAN_SPECIAL_VARS = {"cappi", "hsr", "hsp"} def __init__(self, config: dict[str, Any]): config = dict(config) if "invalid_fill" not in config: stride = config.get("nan_fill_stride") config["invalid_fill"] = { "default": {"method": "scene_mean", "stride": stride}, "variables": { "cappi": {"method": "constant", "value": config.get("nan_special_fill", -250.0)}, "hsr": {"method": "constant", "value": config.get("nan_special_fill", -250.0)}, "hsp": {"method": "constant", "value": config.get("nan_special_fill", -250.0)}, }, } super().__init__(config) self.raw_vars = list(self.config.get("raw_vars", self.config.get("selected_vars", []))) self.physics_formulas = self._load_physics_formulas() self._compiled_physics = [compile(formula, "", "eval") for formula in self.physics_formulas] @property def channels(self) -> list[str]: return list(self.raw_vars) + list(self.physics_formulas) def _load_physics_formulas(self) -> list[str]: formulas = list(self.config.get("physics_formulas", [])) formula_path = self.config.get("physics_formula_path") if formula_path: with Path(formula_path).open("r", encoding="utf-8") as f: formulas.extend(line.strip() for line in f if line.strip()) return formulas def path_for_time(self, timestamp: str): ts = format_time(timestamp) day = ts[:8] tried = [] for root in self.roots: for suffix in (".npy", ".nc"): path = root / day / f"concat_gk2a_radar_{ts}{suffix}" tried.append(path) if path.exists(): return path return None def load_raw_dict(self, timestamp: str) -> dict[str, np.ndarray]: path = self.path_for_time(timestamp) if path is None: raise FileNotFoundError(f"concat file not found for {format_time(timestamp)}") if path.suffix == ".npy": data = np.load(path, allow_pickle=True).item() if not isinstance(data, dict): raise ValueError(f"concat npy must contain dict: {path}") return {k: np.asarray(v, dtype=np.float32) for k, v in data.items()} with xr.open_dataset(path, **self._xr_open_kwargs) as ds: return {k: np.asarray(ds[k].values, dtype=np.float32) for k in ds.data_vars} def get_variable(self, data: dict[str, np.ndarray], name: str) -> np.ndarray: if name in data: return np.asarray(data[name], dtype=np.float32) if name in self.physics_formulas: return self._eval_formula(data, name) raise KeyError(f"{name!r} not found. available={sorted(data)}") def _eval_formula(self, data: dict[str, np.ndarray], formula: str) -> np.ndarray: variables = set(re.findall(r"\b[A-Za-z_]\w*\b", formula)) variables -= {"np", "math", "sin", "cos", "tan", "exp", "log", "sqrt"} missing = sorted(v for v in variables if v not in data) if missing: raise KeyError(f"formula {formula!r} requires missing variables: {missing}") local_dict = {name: np.asarray(data[name], dtype=np.float32) for name in variables} local_dict.update( { "np": np, "math": math, "sin": np.sin, "cos": np.cos, "tan": np.tan, "exp": np.exp, "log": np.log, "sqrt": np.sqrt, } ) code = self._compiled_physics[self.physics_formulas.index(formula)] result = eval(code, {"__builtins__": {}}, local_dict) return np.asarray(result, dtype=np.float32) class ConcatVariableInput(ConcatInput): source_name = "concat_variable" file_dtype = "float16" def __init__(self, config: dict[str, Any]): config = dict(config) var_name = str(config.get("var_name", config.get("variable", ""))).strip() if not var_name: raise ValueError("concat_variable input requires var_name") config["raw_vars"] = [var_name] config["selected_vars"] = [var_name] config.pop("physics_formula_path", None) config.pop("physics_formulas", None) self.var_name = var_name super().__init__(config) @property def channels(self) -> list[str]: return [self.var_name] class L2AIIInput(BaseInput): source_name = "l2_aii" file_dtype = "float16" CAPE_MASK_NAME = "CAPE_mask" def __init__(self, config: dict[str, Any]): config = dict(config) config.setdefault("invalid_fill", {"default": {"method": "zero"}}) self.add_cape_mask = bool(config.get("add_cape_mask", False)) super().__init__(config) @property def channels(self) -> list[str]: channels = list(self.selected_vars) if self.add_cape_mask and self.CAPE_MASK_NAME not in channels: channels.append(self.CAPE_MASK_NAME) return channels def should_normalize(self, name: str) -> bool: return name != self.CAPE_MASK_NAME def should_fill_invalid(self, name: str) -> bool: return name != self.CAPE_MASK_NAME def path_for_time(self, timestamp: str): ts = format_time(timestamp) directory = self.roots[0] / ts[:8] for filename in ( f"l2_aii_{ts}.npy", f"gk2a_ami_le2_aii_ea060lc_{ts}.npy", ): path = directory / filename if path.exists(): return path return None def load_raw_dict(self, timestamp: str) -> dict[str, np.ndarray]: path = self.path_for_time(timestamp) if path is None: raise FileNotFoundError(f"L2 AII file not found for {format_time(timestamp)}") data = np.load(path, allow_pickle=True) if getattr(data, "shape", None) == (): data = data.item() if not isinstance(data, dict): raise ValueError(f"L2 AII npy must contain dict: {path}") return {k: np.asarray(v, dtype=np.float32) for k, v in data.items()} def get_variable(self, data: dict[str, np.ndarray], name: str) -> np.ndarray: if name == self.CAPE_MASK_NAME: if "CAPE" not in data: raise KeyError(f"'CAPE' not found. available={sorted(data)}") return np.isfinite(np.asarray(data["CAPE"], dtype=np.float32)).astype(np.float32) return super().get_variable(data, name)