File size: 10,486 Bytes
76d61a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
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)