Download code/training/src/data_pipeline/input_data.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 17.9 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/data_pipeline/input_data.py
- Command line
-
hf download hf://lsh9034/ci-net/code/training/src/data_pipeline/input_data.py
-
curl -L -o input_data.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/training/src/data_pipeline/input_data.py
17.9 kB
| 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")) | |
| def channels(self) -> list[str]: | |
| return list(self.selected_vars) | |
| 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, "<physics_formula>", "eval") for formula in self.physics_formulas] | |
| 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) | |
| 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) | |
| 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) | |