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