File size: 6,036 Bytes
1aeffbb
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import argparse
import sys
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))

import math

import matplotlib.pyplot as plt
import numpy as np
import torch
import yaml


def load_field(path, time_index=0):
    """Return a prediction as [C, H, W] from common model output layouts."""
    value = torch.load(path, map_location="cpu", weights_only=False)
    if isinstance(value, dict):
        for key in ("prediction", "predictions", "output", "outputs"):
            if key in value:
                value = value[key]
                break
    field = torch.as_tensor(value).detach().cpu().float().numpy()
    if field.ndim == 5:  # [B, T, C, H, W]
        field = field[0, time_index]
    elif field.ndim == 4:  # [B, C, H, W] or [T, C, H, W]
        field = field[0 if field.shape[0] == 1 else time_index]
    elif field.ndim != 3:
        raise ValueError(f"Expected [C,H,W], [B,C,H,W], or [B,T,C,H,W], got {field.shape}")
    if field.ndim != 3:
        raise ValueError(f"Selected output is not [C,H,W]: {field.shape}")
    return field


def load_channel_names(config_path, channel_count):
    if config_path is None:
        return [f"channel_{index}" for index in range(channel_count)]
    with open(config_path, encoding="utf-8") as handle:
        config = yaml.safe_load(handle)
    names = config.get("data", {}).get("channels", [])
    if len(names) != channel_count:
        return [f"channel_{index}" for index in range(channel_count)]
    return names


def choose_channels(names, requested, max_panels):
    if requested:
        selected = []
        for item in requested:
            if item.isdigit():
                index = int(item)
                if not 0 <= index < len(names):
                    raise ValueError(f"Channel index out of range: {index}")
            else:
                if item not in names:
                    raise ValueError(f"Unknown channel: {item}")
                index = names.index(item)
            if index not in selected:
                selected.append(index)
        return selected
    return list(range(min(max_panels, len(names))))


def is_signed_channel(name):
    return name.startswith(("u_", "v_")) or name.startswith("ssh_")


def plot_fields(field, names, indices, output, title, reference=None):
    columns = min(3, len(indices))
    rows = math.ceil(len(indices) / columns)
    has_reference = reference is not None
    fig, axes = plt.subplots(rows, columns, figsize=(5.6 * columns, 4.4 * rows), squeeze=False)
    axes = axes.ravel()
    height, width = field.shape[-2:]
    longitude = np.linspace(0, 360, width, endpoint=False)
    latitude = np.linspace(90, -90, height)
    extent = [longitude[0], longitude[-1], latitude[-1], latitude[0]]

    for axis, index in zip(axes, indices):
        data = field[index]
        reference_data = reference[index] if has_reference else None
        if reference_data is not None:
            data_min = min(np.nanpercentile(data, 2), np.nanpercentile(reference_data, 2))
            data_max = max(np.nanpercentile(data, 98), np.nanpercentile(reference_data, 98))
        else:
            data_min, data_max = np.nanpercentile(data, [2, 98])
        if np.isclose(data_min, data_max):
            data_min, data_max = float(np.nanmin(data)), float(np.nanmax(data) + 1e-6)
        cmap = "RdBu_r" if is_signed_channel(names[index]) else "viridis"
        image = axis.imshow(data, extent=extent, origin="upper", cmap=cmap,
                            vmin=data_min, vmax=data_max, aspect="auto")
        axis.set_title(names[index], fontsize=11, fontweight="bold")
        axis.set_xlabel("Longitude (degrees)")
        axis.set_ylabel("Latitude (degrees)")
        axis.set_xticks([0, 90, 180, 270, 360])
        axis.set_yticks([-90, -45, 0, 45, 90])
        axis.grid(color="white", linewidth=0.35, alpha=0.35)
        colorbar = fig.colorbar(image, ax=axis, fraction=0.046, pad=0.04)
        colorbar.ax.tick_params(labelsize=8)
        stats = f"min {np.nanmin(data):.3g} | max {np.nanmax(data):.3g} | mean {np.nanmean(data):.3g}"
        if reference_data is not None:
            rmse = np.sqrt(np.nanmean((data - reference_data) ** 2))
            stats += f" | RMSE {rmse:.3g}"
        axis.text(0.02, 0.02, stats, transform=axis.transAxes, fontsize=8,
                  color="white", bbox={"facecolor": "black", "alpha": 0.55, "pad": 3})

    for axis in axes[len(indices):]:
        axis.remove()
    fig.suptitle(title, fontsize=15, fontweight="bold")
    fig.tight_layout()
    fig.savefig(output, dpi=180, bbox_inches="tight")
    plt.close(fig)


def main():
    parser = argparse.ArgumentParser()
    parser.add_argument("--input", default=str(ROOT / "result/glonet/data/prediction.pt"))
    parser.add_argument("--output", default=str(ROOT / "result/glonet/prediction.png"))
    parser.add_argument("--config", default=str(ROOT / "conf/config.yaml"))
    parser.add_argument("--reference", default=None, help="Optional .pt truth field for RMSE comparison")
    parser.add_argument("--channel", action="append", help="Channel name or zero-based index; repeatable")
    parser.add_argument("--max-panels", type=int, default=6)
    parser.add_argument("--time-index", type=int, default=0)
    args = parser.parse_args()
    prediction = load_field(args.input, args.time_index)
    names = load_channel_names(args.config, prediction.shape[0])
    indices = choose_channels(names, args.channel, args.max_panels)
    reference = load_field(args.reference, args.time_index) if args.reference else None
    if reference is not None and reference.shape != prediction.shape:
        raise ValueError(f"Prediction/reference shape mismatch: {prediction.shape} vs {reference.shape}")
    Path(args.output).parent.mkdir(parents=True, exist_ok=True)
    plot_fields(prediction, names, indices, args.output,
                f"GLONET ocean forecast | {len(indices)} channel(s)", reference)
    print(f"saved={args.output}")


if __name__ == "__main__":
    main()