File size: 6,943 Bytes
fa2b79f
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
"""Compute paper verification metrics and plot performance and reliability diagrams."""

import json
from pathlib import Path

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import yaml
from sklearn.metrics import average_precision_score, roc_auc_score


ROOT = Path(__file__).resolve().parents[1]


def curves(probability, target, threshold_count):
    thresholds = np.linspace(0, 1, threshold_count)
    pod, pofd, sr, far, csi, bias = [], [], [], [], [], []
    for threshold in thresholds:
        forecast = probability >= threshold
        event = target == 1
        hits = np.sum(forecast & event); false_alarms = np.sum(forecast & ~event)
        misses = np.sum(~forecast & event); negatives = np.sum(~forecast & ~event)
        pod.append(hits / max(hits + misses, 1)); pofd.append(false_alarms / max(false_alarms + negatives, 1))
        sr.append(hits / max(hits + false_alarms, 1)); far.append(false_alarms / max(hits + false_alarms, 1))
        csi.append(hits / max(hits + false_alarms + misses, 1)); bias.append((hits + false_alarms) / max(hits + misses, 1))
    return {key: np.asarray(value) for key, value in (("threshold", thresholds), ("pod", pod), ("pofd", pofd),
            ("sr", sr), ("far", far), ("csi", csi), ("bias", bias))}


def reliability(probability, target, bins):
    edges = np.linspace(0, 1, bins + 1)
    index = np.minimum(np.digitize(probability, edges[1:-1]), bins - 1)
    records, component = [], 0.0
    for bin_index in range(bins):
        mask = index == bin_index
        if not mask.any():
            continue
        forecast_mean, observed_frequency = float(probability[mask].mean()), float(target[mask].mean())
        component += mask.mean() * (forecast_mean - observed_frequency) ** 2
        records.append({"count": int(mask.sum()), "forecast_probability": forecast_mean,
                        "observed_frequency": observed_frequency})
    return records, float(component)


def metrics(probability, target, threshold_count, bins):
    curve = curves(probability, target, threshold_count)
    prevalence = float(target.mean())
    auc = float(roc_auc_score(target, probability)) if 0 < target.sum() < len(target) else 0.5
    aupdc = float(average_precision_score(target, probability)) if target.sum() else 0.0
    positives, negatives = int(target.sum()), int(len(target) - target.sum())
    minimum_aupdc = float(np.mean(np.arange(1, positives + 1) / (np.arange(1, positives + 1) + negatives))) if positives else 0.0
    best = int(np.argmax(curve["csi"])); max_csi = float(curve["csi"][best])
    brier = float(np.mean((probability - target) ** 2))
    reference = prevalence * (1 - prevalence)
    records, reliability_component = reliability(probability, target, bins)
    return {"samples": len(target), "event_rate": prevalence, "auc": auc, "aupdc": aupdc,
            "minimum_aupdc": minimum_aupdc, "naupdc": (aupdc - minimum_aupdc) / max(1 - minimum_aupdc, 1e-12),
            "max_csi": max_csi, "ncsi": (max_csi - prevalence) / max(1 - prevalence, 1e-12),
            "max_csi_threshold": float(curve["threshold"][best]), "pod": float(curve["pod"][best]),
            "pofd": float(curve["pofd"][best]), "sr": float(curve["sr"][best]), "far": float(curve["far"][best]),
            "csi": max_csi, "bias": float(curve["bias"][best]), "brier_score": brier,
            "brier_skill_score": 1 - brier / reference if reference > 0 else 0.0,
            "reliability_component": reliability_component, "reliability_bins": records}, curve


def main():
    config = yaml.safe_load((ROOT / "conf/config.yaml").read_text())
    data = np.load(ROOT / config["paths"]["inference_dir"] / "predictions.npz")
    if str(data["format_version"]) != config["data"]["format_version"]:
        raise ValueError("incompatible prediction format")
    probabilities, targets, groups = data["probabilities"], data["targets"], data["lead_group"]
    if probabilities.shape != targets.shape or probabilities.shape[1:] != (3,):
        raise ValueError("probabilities and targets must have shape [N,3]")
    hazards, group_names = data["hazards"].tolist(), data["lead_group_names"].tolist()
    report, all_curves = {}, {}
    for group_index, group_name in enumerate(group_names):
        report[group_name] = {}
        mask = groups == group_index
        for hazard_index, hazard in enumerate(hazards):
            result, curve = metrics(probabilities[mask, hazard_index], targets[mask, hazard_index],
                                    int(config["evaluation"]["threshold_count"]), int(config["evaluation"]["probability_bins"]))
            report[group_name][hazard] = result; all_curves[(group_name, hazard)] = curve
    numeric = [value for group in report.values() for hazard in group.values() for value in hazard.values()
               if isinstance(value, (int, float))]
    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({"by_lead_group_and_hazard": report,
        "metric_protocol": "Flora et al. AUC, weighted-average-precision AUPDC, minimum-AUPDC NAUPDC, climatology-normalized max CSI, and Brier verification"}, indent=2) + "\n")
    colors = {"tornado": "#b3261e", "hail": "#2e7d32", "wind": "#1565c0"}
    figure, axes = plt.subplots(1, 2, figsize=(10, 4.5), sharex=True, sharey=True)
    for axis, group_name in zip(axes, group_names):
        for hazard in hazards:
            curve = all_curves[(group_name, hazard)]
            axis.plot(curve["sr"], curve["pod"], color=colors[hazard], label=hazard)
        axis.set(xlim=(0, 1), ylim=(0, 1), xlabel="Success ratio", ylabel="Probability of detection", title=group_name.replace("_", " ").title())
        axis.grid(alpha=0.25); axis.legend()
    figure.tight_layout(); figure.savefig(output / "performance.png", dpi=160); plt.close(figure)
    figure, axes = plt.subplots(1, 2, figsize=(10, 4.5), sharex=True, sharey=True)
    for axis, group_name in zip(axes, group_names):
        axis.plot((0, 1), (0, 1), "k--", linewidth=1, label="perfect")
        for hazard in hazards:
            bins = report[group_name][hazard]["reliability_bins"]
            axis.plot([item["forecast_probability"] for item in bins], [item["observed_frequency"] for item in bins],
                      marker="o", color=colors[hazard], label=hazard)
        axis.set(xlim=(0, 1), ylim=(0, 1), xlabel="Forecast probability", ylabel="Observed frequency", title=group_name.replace("_", " ").title())
        axis.grid(alpha=0.25); axis.legend()
    figure.tight_layout(); figure.savefig(output / "reliability.png", dpi=160); plt.close(figure)
    print(f"evaluation={output.relative_to(ROOT)} groups={len(group_names)} hazards={len(hazards)}")


if __name__ == "__main__":
    main()