ci-net / code /training /src /data_pipeline /input_data.py
lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
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"))
@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, "<physics_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)