File size: 5,065 Bytes
eca4864
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
from __future__ import annotations

import argparse
import json
import sys
from pathlib import Path

SCRIPT_DIR = Path(__file__).resolve().parent
if str(SCRIPT_DIR) not in sys.path:
    sys.path.insert(0, str(SCRIPT_DIR))

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np

from common import DEFAULT_CONFIG, load_config, resolve_path


def latitude_weights(height: int) -> np.ndarray:
    latitude = np.linspace(np.pi / 2, -np.pi / 2, height)
    weights = np.cos(latitude).clip(min=0)
    return weights / weights.mean()


def compute_metrics(prediction: np.ndarray, target: np.ndarray) -> dict[str, list[float]]:
    weights = latitude_weights(target.shape[-2])[None, None, None, :, None]
    error = prediction - target
    rmse = np.sqrt(np.mean(error**2 * weights, axis=(0, 1, 3, 4)))
    spatial_weight = weights / (
        weights.sum(axis=(-2, -1), keepdims=True) * target.shape[-1]
    )
    pred_mean = np.sum(prediction * spatial_weight, axis=(-2, -1), keepdims=True)
    target_mean = np.sum(target * spatial_weight, axis=(-2, -1), keepdims=True)
    pred_anomaly = prediction - pred_mean
    target_anomaly = target - target_mean
    numerator = np.sum(pred_anomaly * target_anomaly * weights, axis=(0, 1, 3, 4))
    denominator = np.sqrt(
        np.sum(pred_anomaly**2 * weights, axis=(0, 1, 3, 4))
        * np.sum(target_anomaly**2 * weights, axis=(0, 1, 3, 4))
    )
    acc = numerator / np.maximum(denominator, 1e-12)
    return {"rmse": rmse.tolist(), "acc": acc.tolist()}


def plot_sample(
    prediction: np.ndarray,
    target: np.ndarray,
    variable: str,
    channel_index: int,
    cmap: str,
    output_path: Path,
) -> None:
    predicted = prediction[0, 0, channel_index]
    expected = target[0, 0, channel_index]
    error = predicted - expected
    value_min = min(predicted.min(), expected.min())
    value_max = max(predicted.max(), expected.max())
    error_limit = max(abs(error.min()), abs(error.max()), 1e-12)
    extent = (0, 360, -90, 90)

    figure, axes = plt.subplots(3, 1, figsize=(12, 10), constrained_layout=True)
    image = axes[0].imshow(
        expected, origin="upper", extent=extent, aspect="auto", cmap=cmap,
        vmin=value_min, vmax=value_max,
    )
    axes[0].set_title(f"Target {variable}")
    figure.colorbar(image, ax=axes[0], orientation="vertical")
    image = axes[1].imshow(
        predicted, origin="upper", extent=extent, aspect="auto", cmap=cmap,
        vmin=value_min, vmax=value_max,
    )
    axes[1].set_title(f"Prediction {variable}")
    figure.colorbar(image, ax=axes[1], orientation="vertical")
    image = axes[2].imshow(
        error, origin="upper", extent=extent, aspect="auto", cmap="RdBu_r",
        vmin=-error_limit, vmax=error_limit,
    )
    axes[2].set_title(f"Error {variable}")
    figure.colorbar(image, ax=axes[2], orientation="vertical")
    for axis in axes:
        axis.set_xlabel("Longitude")
        axis.set_ylabel("Latitude")
    output_path.parent.mkdir(parents=True, exist_ok=True)
    figure.savefig(output_path, dpi=160)
    plt.close(figure)


def main() -> None:
    parser = argparse.ArgumentParser(description="Evaluate and plot FCNv2 output")
    parser.add_argument("--config", default=str(DEFAULT_CONFIG))
    parser.add_argument("--input")
    args = parser.parse_args()

    config = load_config(args.config)
    inference_dir = resolve_path(config, config["inference"]["output_dir"])
    input_path = Path(args.input).expanduser().resolve() if args.input else None
    files = [input_path] if input_path else sorted(inference_dir.glob("sample_*.npz"))
    if not files:
        raise FileNotFoundError(f"No inference outputs found in {inference_dir}")

    predictions = []
    targets = []
    for path in files:
        with np.load(path) as data:
            predictions.append(data["prediction"])
            targets.append(data["target"])
    prediction = np.concatenate(predictions)
    target = np.concatenate(targets)
    metrics = compute_metrics(prediction, target)

    output_dir = resolve_path(config, config["visualization"]["output_dir"])
    output_dir.mkdir(parents=True, exist_ok=True)
    (output_dir / "metrics.json").write_text(
        json.dumps(metrics, indent=2), encoding="utf-8"
    )
    variable = config["visualization"]["variable"]
    channel_index = config["data"]["variables"].index(variable)
    sample_index = config["visualization"]["sample_index"]
    if not 0 <= sample_index < prediction.shape[0]:
        raise IndexError(
            f"visualization.sample_index={sample_index} is outside "
            f"the available range [0, {prediction.shape[0] - 1}]"
        )
    plot_sample(
        prediction[sample_index : sample_index + 1],
        target[sample_index : sample_index + 1],
        variable,
        channel_index,
        config["visualization"]["cmap"],
        output_dir / f"{variable}_forecast.png",
    )
    print(output_dir / "metrics.json")
    print(output_dir / f"{variable}_forecast.png")


if __name__ == "__main__":
    main()