ClimateBench / scripts /result.py
zhangrenchao's picture
Publish ClimateBench engineering reproduction
2e913c2 verified
Raw
History Blame Contribute Delete
4.19 kB
"""Compute ClimateBench spatial/global NRMSE and plot all four targets."""
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 area_mean(field: np.ndarray, latitude: np.ndarray) -> np.ndarray:
weights = np.cos(np.deg2rad(latitude)).clip(0)[:, None]
return np.sum(field * weights, axis=(-2, -1)) / (weights.sum() * field.shape[-1])
def climatebench_metrics(prediction: np.ndarray, target: np.ndarray, latitude: np.ndarray, alpha: float) -> dict:
normalization = float(area_mean(target, latitude).mean())
if abs(normalization) <= 1e-12:
raise ValueError("area-weighted target mean must be nonzero for NRMSE normalization")
temporal_rmse = np.sqrt(np.mean((prediction - target) ** 2, axis=0))
spatial = float(area_mean(temporal_rmse, latitude) / normalization)
predicted_global = area_mean(prediction, latitude)
target_global = area_mean(target, latitude)
global_nrmse = float(np.sqrt(np.mean((predicted_global - target_global) ** 2)) / normalization)
return {"spatial_nrmse": spatial, "global_nrmse": global_nrmse,
"total_nrmse": spatial + alpha * global_nrmse}
def main() -> None:
config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
data = np.load(ROOT / config["paths"]["inference"])
if str(data["format_version"]) != config["data"]["format_version"]:
raise ValueError("incompatible prediction format")
start, end = int(data["evaluation_start_year"]), int(data["evaluation_end_year"])
if (start, end) != (int(config["evaluation"]["start_year"]), int(config["evaluation"]["end_year"])):
raise ValueError("evaluation period metadata must be 2080-2100")
names = data["target_names"].tolist()
alpha = float(config["evaluation"]["global_weight"])
metrics = {name: climatebench_metrics(data["predictions"][:, index], data["targets"][:, index],
data["latitude"], alpha) for index, name in enumerate(names)}
report = {"variables": metrics, "mean_total_nrmse": float(np.mean([value["total_nrmse"] for value in metrics.values()])),
"evaluation_protocol": {"scenario": str(data["scenario"]), "period": f"{start}-{end}",
"target_aggregation": str(data["target_aggregation"]),
"spatial_metric": "cosine-area mean of grid-cell temporal RMSE",
"global_weighting": "cosine latitude cell area",
"total": "spatial_nrmse + 5 * global_nrmse",
"normalization": "cosine-area and sample mean of target field"}}
numeric = [number for values in metrics.values() for number in values.values()]
if not np.isfinite(numeric).all():
raise FloatingPointError("evaluation contains NaN or Inf")
output = ROOT / config["paths"]["evaluation_dir"]
output.mkdir(parents=True, exist_ok=True)
(output / "metrics.json").write_text(json.dumps(report, indent=2) + "\n")
figure, axes = plt.subplots(4, 3, figsize=(12, 12), constrained_layout=True)
longitude, latitude = data["longitude"], data["latitude"]
for index, name in enumerate(names):
fields = (data["targets"][:, index].mean(0), data["predictions"][:, index].mean(0),
(data["predictions"][:, index] - data["targets"][:, index]).mean(0))
for column, (field, title) in enumerate(zip(fields, ("Target", "Prediction", "Error"))):
image = axes[index, column].pcolormesh(longitude, latitude, field, shading="auto", cmap="coolwarm")
axes[index, column].set(title=f"{name}: {title}", xlabel="Longitude", ylabel="Latitude")
figure.colorbar(image, ax=axes[index, column], shrink=0.75)
figure.suptitle(f"ClimateBench {str(data['scenario'])} {start}-{end}")
figure.savefig(output / "four_targets.png", dpi=150)
plt.close(figure)
print(f"evaluation={output.relative_to(ROOT)} variables={','.join(names)}")
if __name__ == "__main__":
main()