OneScience-Group/ML-MODIS
05
1"""Pure NumPy bootstrap random-forest regression used by ML-MODIS."""2 3from __future__ import annotations4 5from dataclasses import dataclass6from typing import Any, Dict, List, Optional, Sequence, Tuple7 8import numpy as np9 10 11TARGETS = ("Nd", "reff", "LWP", "CF")12PRESSURE_VARIABLES = ("temperature", "specific_humidity", "relative_humidity", "u_wind", "v_wind", "omega", "geopotential", "cloud_liquid", "cloud_fraction")13PRESSURE_LEVELS = (1000, 950, 900, 850, 800, 750, 700, 650, 600, 550)14SINGLE_FEATURES = (15 "sst", "surface_pressure", "mslp", "skin_temperature", "t2m", "d2m",16 "u10", "v10", "surface_solar_radiation", "surface_thermal_radiation",17 "latent_heat_flux", "sensible_heat_flux", "boundary_layer_height",18 "total_column_water_vapour", "total_column_cloud_liquid", "cape", "cin",19 "low_cloud_cover", "sea_ice_fraction", "precipitation", "cos_sza",20 "latitude", "longitude", "platform_hour",21)22 23 24def feature_names() -> List[str]:25 names = [f"{variable}_{level}hPa" for variable in PRESSURE_VARIABLES for level in PRESSURE_LEVELS]26 names.extend(SINGLE_FEATURES)27 if len(names) != 114:28 raise RuntimeError("The ERA5 predictor ledger must contain exactly 114 features")29 return names30 31 32def regression_metrics(y_true: np.ndarray, y_pred: np.ndarray) -> Dict[str, float]:33 mask = np.isfinite(y_true) & np.isfinite(y_pred)34 if mask.sum() < 2:35 return {"n": int(mask.sum()), "mse": float("nan"), "r2": float("nan"), "pearson": float("nan")}36 y = np.asarray(y_true[mask], dtype=np.float64)37 p = np.asarray(y_pred[mask], dtype=np.float64)38 mse = float(np.mean((y - p) ** 2))39 variance = float(np.sum((y - y.mean()) ** 2))40 r2 = float(1.0 - np.sum((y - p) ** 2) / variance) if variance > 0 else float("nan")41 pearson = float(np.corrcoef(y, p)[0, 1]) if np.std(y) > 0 and np.std(p) > 0 else float("nan")42 return {"n": int(mask.sum()), "mse": mse, "r2": r2, "pearson": pearson}43 44 45@dataclass46class TreeConfig:47 min_leaf: int = 748 max_features: int = 3849 max_depth: Optional[int] = None50 split_candidates: int = 1251 52 53class RandomRegressionTree:54 """CART regressor with random feature subsets and compact array state."""55 56 def __init__(self, config: TreeConfig, seed: int):57 self.config = config58 self.seed = int(seed)59 self.feature: List[int] = []60 self.threshold: List[float] = []61 self.left: List[int] = []62 self.right: List[int] = []63 self.value: List[float] = []64 65 def fit(self, x: np.ndarray, y: np.ndarray) -> "RandomRegressionTree":66 x = np.asarray(x, dtype=np.float32)67 y = np.asarray(y, dtype=np.float64)68 rng = np.random.default_rng(self.seed)69 70 def build(indices: np.ndarray, depth: int) -> int:71 node = len(self.value)72 self.feature.append(-1)73 self.threshold.append(np.nan)74 self.left.append(-1)75 self.right.append(-1)76 self.value.append(float(y[indices].mean()))77 if indices.size < 2 * self.config.min_leaf:78 return node79 if self.config.max_depth is not None and depth >= self.config.max_depth:80 return node81 parent_sse = float(np.sum((y[indices] - y[indices].mean()) ** 2))82 if parent_sse <= 1e-12:83 return node84 n_features = min(self.config.max_features, x.shape[1])85 candidates = rng.choice(x.shape[1], size=n_features, replace=False)86 best: Optional[Tuple[float, int, float, np.ndarray]] = None87 quantiles = np.linspace(0.05, 0.95, self.config.split_candidates)88 for feature in candidates:89 values = x[indices, feature]90 thresholds = np.unique(np.quantile(values, quantiles))91 for threshold in thresholds:92 is_left = values <= threshold93 nl = int(is_left.sum())94 nr = indices.size - nl95 if nl < self.config.min_leaf or nr < self.config.min_leaf:96 continue97 yl, yr = y[indices[is_left]], y[indices[~is_left]]98 score = float(np.sum((yl - yl.mean()) ** 2) + np.sum((yr - yr.mean()) ** 2))99 if best is None or score < best[0]:100 best = (score, int(feature), float(threshold), is_left.copy())101 if best is None or best[0] >= parent_sse - 1e-12:102 return node103 _, split_feature, split_threshold, is_left = best104 self.feature[node] = split_feature105 self.threshold[node] = split_threshold106 self.left[node] = build(indices[is_left], depth + 1)107 self.right[node] = build(indices[~is_left], depth + 1)108 return node109 110 build(np.arange(y.size, dtype=np.int64), 0)111 return self112 113 def predict(self, x: np.ndarray) -> np.ndarray:114 x = np.asarray(x, dtype=np.float32)115 output = np.empty(x.shape[0], dtype=np.float32)116 for row in range(x.shape[0]):117 node = 0118 while self.feature[node] >= 0:119 node = self.left[node] if x[row, self.feature[node]] <= self.threshold[node] else self.right[node]120 output[row] = self.value[node]121 return output122 123 def state_dict(self) -> Dict[str, Any]:124 return {125 "seed": self.seed,126 "config": self.config.__dict__.copy(),127 "feature": np.asarray(self.feature, dtype=np.int32),128 "threshold": np.asarray(self.threshold, dtype=np.float32),129 "left": np.asarray(self.left, dtype=np.int32),130 "right": np.asarray(self.right, dtype=np.int32),131 "value": np.asarray(self.value, dtype=np.float32),132 }133 134 @classmethod135 def from_state_dict(cls, state: Dict[str, Any]) -> "RandomRegressionTree":136 tree = cls(TreeConfig(**state["config"]), int(state["seed"]))137 for name in ("feature", "threshold", "left", "right", "value"):138 setattr(tree, name, np.asarray(state[name]).tolist())139 return tree140 141 142class BootstrapRandomForestRegressor:143 """Regression forest with explicit approximately 60% bootstrap and OOB state."""144 145 def __init__(self, n_trees: int = 100, min_leaf: int = 7, max_features: int = 38,146 bootstrap_fraction: float = 0.6, max_depth: Optional[int] = None,147 split_candidates: int = 12, seed: int = 0):148 if n_trees < 1 or min_leaf < 1 or not 0 < bootstrap_fraction <= 1:149 raise ValueError("Invalid forest configuration")150 self.n_trees = int(n_trees)151 self.bootstrap_fraction = float(bootstrap_fraction)152 self.seed = int(seed)153 self.tree_config = TreeConfig(int(min_leaf), int(max_features), max_depth, int(split_candidates))154 self.trees: List[RandomRegressionTree] = []155 self.oob_indices: List[np.ndarray] = []156 157 def fit(self, x: np.ndarray, y: np.ndarray) -> "BootstrapRandomForestRegressor":158 x = np.asarray(x, dtype=np.float32)159 y = np.asarray(y, dtype=np.float32)160 if x.ndim != 2 or x.shape[1] != 114 or y.shape != (x.shape[0],):161 raise ValueError(f"Expected X [N,114] and y [N], got {x.shape} and {y.shape}")162 rng = np.random.default_rng(self.seed)163 draw_size = max(2 * self.tree_config.min_leaf, int(round(self.bootstrap_fraction * x.shape[0])))164 self.trees, self.oob_indices = [], []165 for _ in range(self.n_trees):166 bootstrap = rng.integers(0, x.shape[0], size=draw_size)167 used = np.zeros(x.shape[0], dtype=bool)168 used[np.unique(bootstrap)] = True169 oob = np.flatnonzero(~used)170 tree_seed = int(rng.integers(0, 2**31 - 1))171 self.trees.append(RandomRegressionTree(self.tree_config, tree_seed).fit(x[bootstrap], y[bootstrap]))172 self.oob_indices.append(oob.astype(np.int32))173 return self174 175 def predict_trees(self, x: np.ndarray) -> np.ndarray:176 if not self.trees:177 raise RuntimeError("Forest is not fitted")178 return np.stack([tree.predict(x) for tree in self.trees], axis=1)179 180 def predict(self, x: np.ndarray) -> np.ndarray:181 return self.predict_trees(x).mean(axis=1)182 183 def oob_predict(self, x: np.ndarray) -> Tuple[np.ndarray, np.ndarray]:184 sums = np.zeros(x.shape[0], dtype=np.float64)185 counts = np.zeros(x.shape[0], dtype=np.int32)186 for tree, indices in zip(self.trees, self.oob_indices):187 if indices.size:188 sums[indices] += tree.predict(x[indices])189 counts[indices] += 1190 prediction = np.full(x.shape[0], np.nan, dtype=np.float32)191 valid = counts > 0192 prediction[valid] = (sums[valid] / counts[valid]).astype(np.float32)193 return prediction, counts194 195 def permutation_importance(self, x: np.ndarray, y: np.ndarray, seed: int = 0) -> np.ndarray:196 """Breiman OOB permuted-predictor delta MSE, averaged over eligible trees."""197 rng = np.random.default_rng(seed)198 deltas = np.zeros(x.shape[1], dtype=np.float64)199 counts = np.zeros(x.shape[1], dtype=np.int32)200 for tree, indices in zip(self.trees, self.oob_indices):201 if indices.size < 2:202 continue203 xo = np.asarray(x[indices], dtype=np.float32)204 yo = np.asarray(y[indices], dtype=np.float32)205 baseline = float(np.mean((yo - tree.predict(xo)) ** 2))206 for feature in range(x.shape[1]):207 changed = xo.copy()208 changed[:, feature] = changed[rng.permutation(indices.size), feature]209 deltas[feature] += float(np.mean((yo - tree.predict(changed)) ** 2)) - baseline210 counts[feature] += 1211 return np.divide(deltas, counts, out=np.zeros_like(deltas), where=counts > 0).astype(np.float32)212 213 def state_dict(self) -> Dict[str, Any]:214 return {215 "n_trees": self.n_trees,216 "bootstrap_fraction": self.bootstrap_fraction,217 "seed": self.seed,218 "tree_config": self.tree_config.__dict__.copy(),219 "trees": [tree.state_dict() for tree in self.trees],220 "oob_indices": self.oob_indices,221 }222 223 @classmethod224 def from_state_dict(cls, state: Dict[str, Any]) -> "BootstrapRandomForestRegressor":225 config = state["tree_config"]226 forest = cls(state["n_trees"], config["min_leaf"], config["max_features"],227 state["bootstrap_fraction"], config["max_depth"],228 config["split_candidates"], state["seed"])229 forest.trees = [RandomRegressionTree.from_state_dict(item) for item in state["trees"]]230 forest.oob_indices = [np.asarray(item, dtype=np.int32) for item in state["oob_indices"]]231 return forest232 233 234def validate_multimodal_keys(data: Dict[str, np.ndarray]) -> None:235 required = ("year", "month", "platform", "latitude", "longitude", "X", "Y")236 missing = [key for key in required if key not in data]237 if missing:238 raise ValueError(f"Missing aligned arrays: {missing}")239 n = data["X"].shape[0]240 if data["X"].shape[1] != 114 or data["Y"].shape != (n, 4):241 raise ValueError("Predictors must be [N,114] and targets [N,4]")242 if any(np.asarray(data[key]).shape[0] != n for key in required[:-2]):243 raise ValueError("Year/month/platform/coordinates are not row-aligned")244 keys = list(zip(data["year"].tolist(), data["month"].tolist(), data["platform"].tolist(),245 np.round(data["latitude"], 4).tolist(), np.round(data["longitude"], 4).tolist()))246 if len(set(keys)) != n:247 raise ValueError("Multimodal year-month-platform-latitude-longitude keys are not unique")248 