| 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() |
|
|