RF-ClimParam / scripts /result.py
zhangrenchao's picture
Upload folder using huggingface_hub
c059069 verified
Raw
History Blame Contribute Delete
4.35 kB
"""Evaluate complete spatial predictions independently at every scale."""
import json
from pathlib import Path
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import yaml
ROOT = Path(__file__).resolve().parents[1]
def scores(target, prediction):
residual = target - prediction
rmse = np.sqrt(np.mean(np.square(residual), axis=0))
denominator = np.sum(np.square(target - target.mean(axis=0)), axis=0)
r2 = 1.0 - np.sum(np.square(residual), axis=0) / np.maximum(denominator, 1e-20)
return r2, rmse
def main():
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
data = np.load(ROOT / config["paths"]["inference"])
output = ROOT / config["paths"]["evaluation_dir"]
output.mkdir(parents=True, exist_ok=True)
report = {"format_version": config["data"]["format_version"], "per_scale": {}}
plotted = []
scales = [str(scale) for scale in data["scales"]]
if scales != list(config["data"]["scales"]):
raise ValueError("prediction scales do not match configuration")
for scale_index, scale in enumerate(scales):
report["per_scale"][scale] = {"grid": data[f"{scale}_grid_shape"].tolist()}
for task, names in (("rf_tend", data["tend_output_names"]), ("rf_diff", data["diff_output_names"])):
target = data[f"{scale}_{task[3:]}_targets"]
prediction = data[f"{scale}_{task[3:]}_predictions"]
r2, rmse = scores(target, prediction)
if not np.isfinite(r2).all() or not np.isfinite(rmse).all():
raise FloatingPointError(f"non-finite offline metrics for {scale}/{task}")
report["per_scale"][scale][task] = {
"outputs": [{"name": str(name), "r2": float(r2[i]), "rmse": float(rmse[i])}
for i, name in enumerate(names)],
"mean_r2": float(r2.mean()), "mean_rmse": float(rmse.mean())}
plotted.append((scale, task, float(r2.mean()), float(rmse.mean())))
ny, nx = map(int, config["evaluation"]["online_proxy_grid"])
target = data["online_x32_native_precipitation_target_3h"].reshape(ny, nx).mean(axis=1)
prediction = data["online_x32_native_precipitation_prediction_3h"].reshape(ny, nx).mean(axis=1)
proxy_r2, proxy_rmse = scores(target[:, None], prediction[:, None])
report["online_engineering_proxy"] = {
"scale": "x32", "grid": [ny, nx], "time_steps": 1,
"metric": "zonal-mean 3h precipitation",
"r2": float(proxy_r2[0]), "rmse": float(proxy_rmse[0]),
"target_p99_9_3h": float(np.percentile(data["online_x32_native_precipitation_target_3h"], 99.9)),
"prediction_p99_9_3h": float(np.percentile(data["online_x32_native_precipitation_prediction_3h"], 99.9)),
"note": "Engineering proxy for online coupling, not the paper's 600-day SAM simulation."}
numeric = [scale[task]["mean_r2"] for scale in report["per_scale"].values()
for task in ("rf_tend", "rf_diff")]
numeric += [report["online_engineering_proxy"][key] for key in ("r2", "rmse", "target_p99_9_3h", "prediction_p99_9_3h")]
if not np.isfinite(numeric).all():
raise FloatingPointError("evaluation report contains non-finite values")
(output / "metrics.json").write_text(json.dumps(report, indent=2) + "\n")
figure, axes = plt.subplots(1, 2, figsize=(12, 4.5))
labels = [f"{scale}\n{task[3:]}" for scale, task, _, _ in plotted]
axes[0].bar(labels, [item[2] for item in plotted], color=["#275d6c", "#d98b3a"] * 4)
axes[0].axhline(0, color="black", linewidth=0.7); axes[0].set_ylabel("Mean output R2")
axes[0].tick_params(axis="x", labelsize=8); axes[0].set_title("Complete-grid column prediction")
latitude = np.linspace(-90, 90, ny)
axes[1].plot(latitude, target, label="target", color="#202020", linewidth=2)
axes[1].plot(latitude, prediction, label="RF prediction", color="#c44e52", linewidth=2)
axes[1].set(xlabel="Latitude (degrees)", ylabel="3 h precipitation (proxy units)",
title=f"x32 native zonal mean on {ny}x{nx} grid")
axes[1].legend(); figure.tight_layout(); figure.savefig(output / "comparison.png", dpi=150)
plt.close(figure)
print(f"evaluation={output.relative_to(ROOT)} online_proxy_r2={proxy_r2[0]:.3f}")
if __name__ == "__main__":
main()