Download scripts/result.py from OneScience-Group/RF-ClimParam: direct link, hf CLI and curl.
- Browser
- Download file 4.35 kB
-
https://huggingface.co/OneScience-Group/RF-ClimParam/resolve/main/scripts/result.py
- Command line
-
hf download hf://OneScience-Group/RF-ClimParam/scripts/result.py
-
curl -L -o result.py https://huggingface.co/OneScience-Group/RF-ClimParam/resolve/main/scripts/result.py
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() | |