Team Ai
Modelpublic

OneScience-Group/ML-MODIS

sourceHugging Faceapache-2.0updated 1mo agoView on Hugging Face
0likes5downloads
ml_modis.py248 linesDownload Raw Back to model
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