MeteoNorm-RF / model /meteonorm_rf.py
zhangrenchao's picture
Publish MeteoNorm-RF reproduction
b11ef36 verified
Raw History Blame Contribute Delete
7.75 kB
"""Pure NumPy multi-output random forest for Beijing air-quality normalization."""
from dataclasses import dataclass
import pickle
import numpy as np
FORMAT_VERSION = "meteonorm_rf_v1"
MODEL_NAME = "MeteoNorm-RF"
POLLUTANTS = ("PM2.5", "PM10", "NO2", "SO2", "O3", "CO")
STATIONS = ("dongsi", "tiantan", "guanyuan", "wanshouxigong", "aotizhongxin",
"nongzhanguan", "wanliu", "beibuxinqu", "zhiwuyuan", "fengtaihuayuan",
"yungang", "gucheng")
BASE_FEATURES = ("ttrend", "day_of_year", "weekend", "hour", "wind_speed",
"wind_direction", "pressure", "temperature", "relative_humidity")
FEATURE_NAMES = BASE_FEATURES + tuple(f"station={name}" for name in STATIONS)
def encode_features(data):
"""Encode station as an explicit 12-column one-hot categorical feature."""
station = np.asarray(data["station_id"], dtype=np.int64)
if np.any((station < 0) | (station >= len(STATIONS))):
raise ValueError("station_id must be in [0, 11]")
numeric = np.column_stack([np.asarray(data[name], dtype=np.float32)
for name in BASE_FEATURES])
one_hot = np.eye(len(STATIONS), dtype=np.float32)[station]
result = np.concatenate([numeric, one_hot], axis=1)
if result.shape[1] != len(FEATURE_NAMES) or not np.isfinite(result).all():
raise ValueError("features must be finite and follow the 21-column protocol")
return result.astype(np.float32)
@dataclass
class TreeConfig:
max_depth: int = 8
min_samples_leaf: int = 8
max_features: object = "sqrt"
split_candidates: int = 12
class RegressionTree:
def __init__(self, config, seed=0):
self.config = config
self.rng = np.random.default_rng(seed)
self.nodes = []
def _feature_count(self, total):
value = self.config.max_features
if value == "sqrt":
return max(1, int(np.sqrt(total)))
if isinstance(value, float):
return max(1, min(total, int(np.ceil(total * value))))
return max(1, min(total, int(value)))
def fit(self, x, y):
self.nodes = []
self._grow(np.asarray(x, np.float32), np.asarray(y, np.float32),
np.arange(len(x)), 0)
return self
def _grow(self, x, y, indices, depth):
node_id = len(self.nodes)
self.nodes.append(None)
value = y[indices].mean(0).astype(np.float32)
minimum = int(self.config.min_samples_leaf)
if depth >= self.config.max_depth or len(indices) < 2 * minimum:
self.nodes[node_id] = {"value": value}
return node_id
best = None
features = self.rng.choice(x.shape[1], self._feature_count(x.shape[1]), replace=False)
for feature in features:
values = x[indices, feature]
low, high = float(values.min()), float(values.max())
if low >= high:
continue
for threshold in self.rng.uniform(low, high, self.config.split_candidates):
mask = values <= threshold
left, right = indices[mask], indices[~mask]
if len(left) < minimum or len(right) < minimum:
continue
loss = (np.square(y[left] - y[left].mean(0)).sum() +
np.square(y[right] - y[right].mean(0)).sum())
if best is None or loss < best[0]:
best = (float(loss), int(feature), float(threshold), left, right)
if best is None:
self.nodes[node_id] = {"value": 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):
output = []
for row in np.asarray(x, np.float32):
node = self.nodes[0]
while "value" not in node:
branch = "left" if row[node["feature"]] <= node["threshold"] else "right"
node = self.nodes[node[branch]]
output.append(node["value"])
return np.asarray(output, dtype=np.float32)
class MultiOutputRandomForest:
def __init__(self, n_trees=12, seed=2019, **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, tree_indices=None):
x, y = np.asarray(x, np.float32), np.asarray(y, np.float32)
indices = range(self.n_trees) if tree_indices is None else tree_indices
self.trees = []
for index in indices:
rng = np.random.default_rng(self.seed + 7919 * (index + 1))
sample = rng.integers(0, len(x), len(x))
tree = RegressionTree(self.tree_config, self.seed + 104729 * (index + 1))
self.trees.append((int(index), tree.fit(x[sample], y[sample])))
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": [(index, tree.nodes) for index, tree in self.trees]}
@classmethod
def from_state_dict(cls, state):
model = cls(state["n_trees"], state["seed"], **state["tree_config"])
for index, nodes in state["trees"]:
tree = RegressionTree(model.tree_config)
tree.nodes = nodes
model.trees.append((index, tree))
return model
class MeteoNormRF:
def __init__(self, forest):
self.forest = forest
self.statistics = {}
def fit(self, x, y, tree_indices=None):
x, y = np.asarray(x, np.float32), np.asarray(y, np.float32)
self.statistics = {"x_mean": x.mean(0), "x_std": np.maximum(x.std(0), 1e-6),
"y_mean": y.mean(0), "y_std": np.maximum(y.std(0), 1e-6)}
self.forest.fit((x - self.statistics["x_mean"]) / self.statistics["x_std"],
(y - self.statistics["y_mean"]) / self.statistics["y_std"], tree_indices)
return self
def predict(self, x):
x = np.asarray(x, np.float32)
value = self.forest.predict((x - self.statistics["x_mean"]) / self.statistics["x_std"])
return np.maximum(value * self.statistics["y_std"] + self.statistics["y_mean"], 0.0)
def state_dict(self):
return {"forest": self.forest.state_dict(), "statistics": self.statistics}
@classmethod
def from_state_dict(cls, state):
model = cls(MultiOutputRandomForest.from_state_dict(state["forest"]))
model.statistics = state["statistics"]
return model
def merge_states(states):
base = states[0]
trees = [item for state in states for item in state["forest"]["trees"]]
base["forest"]["trees"] = sorted(trees, key=lambda item: item[0])
return base
def save_checkpoint(path, model, metadata):
with open(path, "wb") as stream:
pickle.dump({"format_version": FORMAT_VERSION, "model_name": MODEL_NAME,
"model": model.state_dict(), "metadata": metadata}, stream)
def load_checkpoint(path):
with open(path, "rb") as stream:
state = pickle.load(stream)
if state.get("format_version") != FORMAT_VERSION or state.get("model_name") != MODEL_NAME:
raise ValueError("incompatible checkpoint")
return MeteoNormRF.from_state_dict(state["model"]), state["metadata"]