lsh9034's picture
Add files using upload-large-folder tool
76d61a0 verified
Raw History Blame Contribute Delete
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)
@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)