RF-ClimParam / model /rf_climparam.py
zhangrenchao's picture
Upload folder using huggingface_hub
c059069 verified
Raw
History Blame Contribute Delete
10.5 kB
"""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()}