ACE2 / scripts /result.py
zhangrenchao's picture
Publish ACE2 reproduction
380b161 verified
Raw
History Blame Contribute Delete
3.16 kB
import json
from pathlib import Path
import sys
import numpy as np
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.ace2 import PRECIPITATION, Q_INDICES, SURFACE_PRESSURE, load_config
def weighted_mean(x, w):
return np.sum(x * w, axis=(-2, -1)) / np.sum(np.broadcast_to(w, x.shape), axis=(-2, -1))
def rmse(pred, target, w):
return float(np.sqrt(np.sum((pred - target) ** 2 * w) / np.sum(np.broadcast_to(w, pred.shape))))
def main():
cfg = load_config(ROOT)
data = np.load(ROOT / cfg["data"]["path"])
steps = cfg["inference"]["steps"]
truth = data["state"][0, 1:steps + 1].astype(np.float32)
pred = np.load(ROOT / cfg["inference"]["output"])["forecast"].astype(np.float32)
initial = data["state"][0, 0].astype(np.float32)
persistence = np.broadcast_to(initial, truth.shape)
w = np.cos(np.deg2rad(np.linspace(-89.5, 89.5, 180, dtype=np.float32)))[None, None, :, None]
pg, tg = weighted_mean(pred, w), weighted_mean(truth, w)
denom = np.sum((tg - tg.mean(axis=0, keepdims=True)) ** 2)
r2 = float(1 - np.sum((pg - tg) ** 2) / max(float(denom), 1e-12))
dry0 = weighted_mean(initial[SURFACE_PRESSURE] - initial[list(Q_INDICES)].sum(0), w[0, 0])
dry = weighted_mean(pred[:, SURFACE_PRESSURE] - pred[:, list(Q_INDICES)].sum(1), w[0, 0])
previous = np.concatenate((initial[None], pred[:-1]), axis=0)
water_previous = weighted_mean(previous[:, list(Q_INDICES)].sum(1), w[0, 0])
water = weighted_mean(pred[:, list(Q_INDICES)].sum(1) + pred[:, PRECIPITATION], w[0, 0])
model_rmse, baseline_rmse = rmse(pred, truth, w), rmse(persistence, truth, w)
metrics = {
"area_weighted_rmse": model_rmse,
"global_mean_r2": r2,
"conservation": {
"max_abs_global_dry_mass_error": float(np.max(np.abs(dry - dry0))),
"max_abs_global_moisture_closure_error": float(np.max(np.abs(water - water_previous))),
},
"comparison": {
"persistence_area_weighted_rmse": baseline_rmse,
"rmse_skill_vs_persistence": float(1 - model_rmse / baseline_rmse),
},
}
output = ROOT / cfg["evaluation"]["output"]
output.parent.mkdir(parents=True, exist_ok=True)
output.write_text(json.dumps(metrics, indent=2) + "\n", encoding="utf-8")
figure = ROOT / cfg["evaluation"]["figure"]
figure.parent.mkdir(parents=True, exist_ok=True)
fig, axes = plt.subplots(1, 2, figsize=(9, 3.8), constrained_layout=True)
axes[0].bar(["ACE2", "Persistence"], [model_rmse, baseline_rmse], color=["#287f71", "#8c96a8"])
axes[0].set(ylabel="Area-weighted RMSE", title="Forecast error")
axes[1].bar(["Dry mass", "Moisture"], [metrics["conservation"]["max_abs_global_dry_mass_error"], metrics["conservation"]["max_abs_global_moisture_closure_error"]], color="#d69b36")
axes[1].set_yscale("log"); axes[1].set(title="Conservation residual", ylabel="Maximum absolute error")
fig.savefig(figure, dpi=150); plt.close(fig)
print(json.dumps(metrics, indent=2))
if __name__ == "__main__":
main()