File size: 2,212 Bytes
7a7efc9
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
"""Evaluate all 569 Jive stations hourly and plot station summaries."""

import json
import sys
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
import yaml

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from model.improver_aifs import crps_from_thresholds


def main():
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    data = np.load(ROOT / config["paths"]["inference"])
    raw, blend, truth = data["raw_aifs"], data["blend_expected"], data["analyses"]
    if raw.shape != (2, 241, 3, 569):
        raise ValueError(f"evaluation requires [2,241,3,569], received {raw.shape}")
    metrics = {"note": "engineering validation, not paper performance", "verification_stations": 569, "hourly_leads": 241, "valid_time_folds": 2, "by_variable": {}}
    for variable, name in enumerate(data["variables"]):
        probability, thresholds = data[f"probability_blend_{name}"], data[f"thresholds_{name}"]
        metrics["by_variable"][str(name)] = {
            "raw_mse": float(np.mean((raw[:, :, variable] - truth[:, :, variable]) ** 2)),
            "blend_mse": float(np.mean((blend[:, :, variable] - truth[:, :, variable]) ** 2)),
            "blend_bias": float(np.mean(blend[:, :, variable] - truth[:, :, variable])),
            "blend_crps": crps_from_thresholds(probability, thresholds, truth[:, :, variable]),
        }
    output = ROOT / config["paths"]["evaluation_dir"]
    output.mkdir(parents=True, exist_ok=True)
    (output / "metrics.json").write_text(json.dumps(metrics, indent=2) + "\n")
    fig, axes = plt.subplots(2, 2, figsize=(11, 7), constrained_layout=True)
    for axis, lead in zip(axes.flat, [0, 24, 120, 240]):
        axis.scatter(data["station_longitude"], data["station_latitude"], c=np.abs(blend[0, lead, 0] - truth[0, lead, 0]), s=8, cmap="viridis")
        axis.set_title(f"569-station temperature absolute error | {lead} h")
        axis.set_xlabel("longitude"); axis.set_ylabel("latitude")
    fig.savefig(output / "multi_lead_temperature.png", dpi=150); plt.close(fig)
    print("saved=result/evaluation stations=569 hourly_leads=241 folds=2 mse_crps_bias=ok")


if __name__ == "__main__":
    main()