Download code/final_preprocess/src/data_pipeline/label.py from lsh9034/ci-net: direct link, hf CLI and curl.
- Browser
- Download file 10.5 kB
-
https://huggingface.co/lsh9034/ci-net/resolve/main/code/final_preprocess/src/data_pipeline/label.py
- Command line
-
hf download hf://lsh9034/ci-net/code/final_preprocess/src/data_pipeline/label.py
-
curl -L -o label.py https://huggingface.co/lsh9034/ci-net/resolve/main/code/final_preprocess/src/data_pipeline/label.py
10.5 kB
| 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) | |
| 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 | |
| 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) | |
| 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) | |
| 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) | |