"""Pure NumPy joint multi-output random forests for the SAM parameterization.""" from dataclasses import dataclass import numpy as np FORMAT_VERSION = "rf_climparam_v2" MODEL_NAME = "RF-ClimParam" SCALES = ("x4", "x8", "x16", "x32") TEND_INPUT_NAMES = tuple([f"T_{i:02d}" for i in range(48)] + [f"qT_{i:02d}" for i in range(48)] + [f"qp_{i:02d}" for i in range(48)] + ["abs_y"]) TEND_OUTPUT_NAMES = tuple([f"hL_tend_{i:02d}" for i in range(48)] + [f"qT_tend_{i:02d}" for i in range(48)] + [f"qp_tend_{i:02d}" for i in range(48)]) DIFF_INPUT_NAMES = tuple([f"T_low_{i:02d}" for i in range(15)] + [f"qT_low_{i:02d}" for i in range(15)] + [f"u_low_{i:02d}" for i in range(15)] + [f"v_nh_low_{i:02d}" for i in range(15)] + ["windsurf", "abs_y"]) DIFF_OUTPUT_NAMES = tuple([f"Dbar_{i:02d}" for i in range(15)] + ["hL_surface_flux", "qT_surface_flux"]) def _check_array(name, value, features): value = np.asarray(value) if value.ndim != 2 or value.shape[1] != features: raise ValueError(f"{name} must have shape [N,{features}], got {value.shape}") if value.dtype not in (np.float32, np.float64) or not np.isfinite(value).all(): raise ValueError(f"{name} must be finite float32/float64") return value.astype(np.float32, copy=False) @dataclass class TreeConfig: max_depth: int = 5 min_samples_leaf: int = 3 max_features: object = "sqrt" split_candidates: int = 8 class ExtraRandomRegressionTree: """Randomized recursive tree whose leaves hold one joint output vector.""" def __init__(self, config, seed=0): self.config = config self.rng = np.random.default_rng(seed) self.nodes = [] def fit(self, x, y): x, y = np.asarray(x, np.float32), np.asarray(y, np.float32) self.nodes = [] self._grow(x, y, np.arange(len(x)), 0) return self def _feature_count(self, total): value = self.config.max_features if value == "sqrt": return max(1, int(np.sqrt(total))) if value == "log2": return max(1, int(np.log2(total))) if isinstance(value, float): return max(1, min(total, int(np.ceil(value * total)))) return max(1, min(total, int(value))) def _grow(self, x, y, indices, depth): node_id = len(self.nodes) self.nodes.append(None) leaf_value = y[indices].mean(axis=0).astype(np.float32) minimum = int(self.config.min_samples_leaf) if depth >= int(self.config.max_depth) or len(indices) < 2 * minimum: self.nodes[node_id] = {"value": leaf_value} return node_id features = self.rng.choice(x.shape[1], self._feature_count(x.shape[1]), replace=False) best = None parent_sse = float(np.square(y[indices] - leaf_value).sum()) for feature in features: values = x[indices, feature] low, high = float(values.min()), float(values.max()) if not low < high: continue thresholds = self.rng.uniform(low, high, int(self.config.split_candidates)) for threshold in thresholds: mask = values <= threshold left, right = indices[mask], indices[~mask] if len(left) < minimum or len(right) < minimum: continue left_mean, right_mean = y[left].mean(0), y[right].mean(0) loss = float(np.square(y[left] - left_mean).sum() + np.square(y[right] - right_mean).sum()) if best is None or loss < best[0]: best = (loss, int(feature), float(threshold), left, right) if best is None or best[0] >= parent_sse - 1e-10: self.nodes[node_id] = {"value": leaf_value} return node_id _, feature, threshold, left, right = best self.nodes[node_id] = {"feature": feature, "threshold": threshold, "left": self._grow(x, y, left, depth + 1), "right": self._grow(x, y, right, depth + 1)} return node_id def predict(self, x): outputs = [] for row in np.asarray(x): node = self.nodes[0] while "value" not in node: node = self.nodes[node["left"] if row[node["feature"]] <= node["threshold"] else node["right"]] outputs.append(node["value"]) return np.asarray(outputs, dtype=np.float32) def state_dict(self): return {"config": vars(self.config), "nodes": self.nodes} @classmethod def from_state_dict(cls, state): tree = cls(TreeConfig(**state["config"])) tree.nodes = state["nodes"] return tree class JointRandomForestRegressor: """Bootstrap ensemble retaining inseparable multi-output leaf predictions.""" def __init__(self, n_trees=2, seed=0, **tree_options): self.n_trees, self.seed = int(n_trees), int(seed) self.tree_config = TreeConfig(**tree_options) self.trees = [] def fit(self, x, y): x, y = np.asarray(x, np.float32), np.asarray(y, np.float32) rng = np.random.default_rng(self.seed) self.trees = [] for index in range(self.n_trees): bootstrap = rng.integers(0, len(x), size=len(x)) tree = ExtraRandomRegressionTree(self.tree_config, self.seed + 1009 * (index + 1)) self.trees.append(tree.fit(x[bootstrap], y[bootstrap])) return self def predict(self, x): if not self.trees: raise RuntimeError("forest is not fitted") return np.mean([tree.predict(x) for tree in self.trees], axis=0, dtype=np.float32) def state_dict(self): return {"n_trees": self.n_trees, "seed": self.seed, "tree_config": vars(self.tree_config), "trees": [tree.state_dict() for tree in self.trees]} @classmethod def from_state_dict(cls, state): forest = cls(state["n_trees"], state["seed"], **state["tree_config"]) forest.trees = [ExtraRandomRegressionTree.from_state_dict(item) for item in state["trees"]] return forest class StandardizedForest: """Block-standardized wrapper; one scalar mean/std is used per variable block.""" def __init__(self, forest, input_slices, output_slices, nonnegative_slice=None): self.forest = forest self.input_slices, self.output_slices = input_slices, output_slices self.nonnegative_slice = nonnegative_slice self.statistics = {} @staticmethod def _statistics(array, slices): means, stds = np.zeros(array.shape[1], np.float32), np.ones(array.shape[1], np.float32) for start, stop in slices: mean = float(array[:, start:stop].mean()) std = max(float(array[:, start:stop].std()), 1e-6) means[start:stop], stds[start:stop] = mean, std return means, stds def fit(self, x, y): x = _check_array("inputs", x, self.input_slices[-1][1]) y = _check_array("targets", y, self.output_slices[-1][1]) x_mean, x_std = self._statistics(x, self.input_slices) y_mean, y_std = self._statistics(y, self.output_slices) self.statistics = {"input_mean": x_mean, "input_std": x_std, "output_mean": y_mean, "output_std": y_std} self.forest.fit((x - x_mean) / x_std, (y - y_mean) / y_std) return self def predict(self, x): x = _check_array("inputs", x, len(self.statistics["input_mean"])) prediction = self.forest.predict((x - self.statistics["input_mean"]) / self.statistics["input_std"]) prediction = prediction * self.statistics["output_std"] + self.statistics["output_mean"] if self.nonnegative_slice is not None: prediction[:, self.nonnegative_slice[0]:self.nonnegative_slice[1]] = np.maximum( prediction[:, self.nonnegative_slice[0]:self.nonnegative_slice[1]], 0.0) return prediction.astype(np.float32) def state_dict(self): return {"forest": self.forest.state_dict(), "input_slices": self.input_slices, "output_slices": self.output_slices, "nonnegative_slice": self.nonnegative_slice, "statistics": self.statistics} @classmethod def from_state_dict(cls, state): model = cls(JointRandomForestRegressor.from_state_dict(state["forest"]), state["input_slices"], state["output_slices"], state["nonnegative_slice"]) model.statistics = state["statistics"] return model def build_pair(config, seed): options = config["engineering"] common = {"n_trees": options["trees"], "max_depth": options["max_depth"], "min_samples_leaf": options["min_samples_leaf"], "max_features": options["max_features"], "split_candidates": options["split_candidates"]} tend = StandardizedForest(JointRandomForestRegressor(seed=seed, **common), [(0, 48), (48, 96), (96, 144), (144, 145)], [(0, 48), (48, 96), (96, 144)]) diff = StandardizedForest(JointRandomForestRegressor(seed=seed + 1, **common), [(0, 15), (15, 30), (30, 45), (45, 60), (60, 61), (61, 62)], [(0, 15), (15, 16), (16, 17)], nonnegative_slice=(0, 15)) return {"rf_tend": tend, "rf_diff": diff} def load_models(checkpoint): if checkpoint.get("format_version") != FORMAT_VERSION or checkpoint.get("model_name") != MODEL_NAME: raise ValueError("incompatible checkpoint model/format_version") if not isinstance(checkpoint.get("model"), dict) or set(checkpoint["model"]) != set(SCALES): raise ValueError("checkpoint model must contain all four scale forest states") expected = {"rf_tend_input": 145, "rf_tend_output": 144, "rf_diff_input": 62, "rf_diff_output": 17} if checkpoint.get("model_config", {}).get("dimensions") != expected: raise ValueError("checkpoint model_config dimensions are incompatible") return {scale: {name: StandardizedForest.from_state_dict(state) for name, state in pair.items()} for scale, pair in checkpoint["model"].items()}