FuXi-S2S / scripts /result.py
yzt15806542928's picture
Upload folder using huggingface_hub
3fef26f verified
Raw
History Blame Contribute Delete
6.97 kB
import argparse
from pathlib import Path
import matplotlib.pyplot as plt
import numpy as np
from config import load_config, project_path
DEFAULT_VARIABLES = ("t2m", "msl", "tp", "tcwv")
VARIABLE_STYLES = {
"t2m": ("2 m temperature", "Temperature (K)", "coolwarm"),
"d2m": ("2 m dew point", "Temperature (K)", "coolwarm"),
"sst": ("Sea-surface temperature", "Temperature (K)", "coolwarm"),
"msl": ("Mean sea-level pressure", "Pressure (model units)", "viridis"),
"tp": ("Total precipitation", "Precipitation (model units)", "Blues"),
"tcwv": ("Total column water vapour", "Water vapour (model units)", "YlGnBu"),
}
def find_input(output_dir):
candidates = sorted(Path(output_dir).glob("*.npy"), key=lambda path: path.stat().st_mtime)
if not candidates:
raise FileNotFoundError(
f"No inference result (*.npy) was found in {output_dir}. "
"Run scripts/inference.py first or pass --input explicitly."
)
return candidates[-1]
def select_forecast(array, step, ensemble):
if array.ndim == 5:
if ensemble >= array.shape[0]:
raise IndexError(f"Ensemble index {ensemble} is outside array shape {array.shape}")
array = array[ensemble]
if array.ndim == 4:
if step >= array.shape[0]:
raise IndexError(f"Forecast step {step} is outside array shape {array.shape}")
array = array[step]
if array.ndim != 3:
raise ValueError(
"Expected an inference result shaped [time, channel, lat, lon], "
f"[channel, lat, lon], or [ensemble, time, channel, lat, lon], got {array.shape}"
)
return np.asarray(array, dtype=np.float32)
def parse_variables(value, variables):
names = [name.strip() for name in value.split(",") if name.strip()]
if not names:
raise ValueError("At least one variable must be selected with --variables")
unknown = [name for name in names if name not in variables]
if unknown:
raise ValueError(f"Unknown variables: {unknown}. Available variables are configured in config.yaml")
return names
def plot_field(axis, colorbar_axis, field, variable, latitudes, longitudes):
finite = field[np.isfinite(field)]
if finite.size == 0:
axis.set_title(f"{VARIABLE_STYLES.get(variable, (variable,))[0]} (no finite data)")
axis.set_facecolor("#d9dde3")
axis.text(
0.5,
0.5,
"No finite values",
transform=axis.transAxes,
ha="center",
va="center",
fontsize=12,
color="#263238",
)
axis.set_xlabel("Longitude (degrees)")
axis.set_ylabel("Latitude (degrees)")
axis.set_xticks(np.arange(0, 361, 60))
axis.set_yticks(np.arange(-90, 91, 30))
colorbar_axis.set_visible(False)
return
low, high = np.percentile(finite, [2, 98])
if low == high:
low, high = float(finite.min()), float(finite.max() + 1.0)
label, colorbar_label, colormap = VARIABLE_STYLES.get(
variable, (variable, "Model output", "viridis")
)
image = axis.imshow(
field,
cmap=colormap,
vmin=low,
vmax=high,
extent=(longitudes[0], longitudes[-1], latitudes[-1], latitudes[0]),
interpolation="nearest",
aspect="auto",
)
axis.set_title(label)
axis.set_xlabel("Longitude (degrees)")
axis.set_ylabel("Latitude (degrees)")
axis.set_xticks(np.arange(0, 361, 60))
axis.set_yticks(np.arange(-90, 91, 30))
axis.grid(color="white", alpha=0.3, linewidth=0.6)
axis.text(
0.02,
0.03,
f"range {finite.min():.3g} to {finite.max():.3g}\n"
f"mean {finite.mean():.3g} | p02-p98 {low:.3g}-{high:.3g}",
transform=axis.transAxes,
color="white",
fontsize=8,
va="bottom",
bbox={"facecolor": "black", "alpha": 0.55, "pad": 3, "edgecolor": "none"},
)
colorbar = plt.colorbar(image, cax=colorbar_axis, orientation="horizontal")
colorbar.set_label(colorbar_label)
def main():
config = load_config()
data_config = config["data"]
inference_config = config["inference"]
variables = data_config["variables"]
parser = argparse.ArgumentParser(
description="Render FuXi-S2S forecast fields as a geographic multi-panel image."
)
parser.add_argument(
"--input",
default=None,
help="Inference .npy file; defaults to the newest file in inference.output_dir",
)
parser.add_argument(
"--output",
default=None,
help="PNG output path; defaults to inference.visualization_dir/<input-stem>.png",
)
parser.add_argument("--title", default=inference_config["visualization_title"])
parser.add_argument("--step", type=int, default=0, help="Forecast time index to display")
parser.add_argument("--ensemble", type=int, default=0, help="Ensemble member to display")
parser.add_argument(
"--variables",
default=",".join(name for name in DEFAULT_VARIABLES if name in variables),
help="Comma-separated variable names to display",
)
args = parser.parse_args()
input_path = Path(args.input) if args.input else find_input(project_path(inference_config["output_dir"]))
array = select_forecast(np.load(input_path), args.step, args.ensemble)
selected_variables = parse_variables(args.variables, variables)
if array.shape[0] != len(variables):
raise ValueError(
f"Inference result has {array.shape[0]} channels, but config.yaml defines {len(variables)} variables"
)
output = Path(args.output) if args.output else project_path(inference_config["visualization_dir"]) / f"{input_path.stem}.png"
output.parent.mkdir(parents=True, exist_ok=True)
height, width = array.shape[-2:]
latitudes = np.linspace(90.0, -90.0, height)
longitudes = np.linspace(0.0, 360.0, width, endpoint=False)
figure = plt.figure(figsize=(16, 9), constrained_layout=True)
layout = figure.add_gridspec(
2,
len(selected_variables),
height_ratios=[12, 1],
width_ratios=[1] * len(selected_variables),
)
figure.suptitle(
f"{args.title}\n{input_path.name} | forecast step {args.step} | ensemble {args.ensemble}",
fontsize=15,
fontweight="bold",
)
for index, variable in enumerate(selected_variables):
axis = figure.add_subplot(layout[0, index])
colorbar_axis = figure.add_subplot(layout[1, index])
variable_index = variables.index(variable)
plot_field(axis, colorbar_axis, array[variable_index], variable, latitudes, longitudes)
colorbar_axis.set_xticks([])
colorbar_axis.set_yticks([])
figure.savefig(output, dpi=180, facecolor="white")
plt.close(figure)
print(f"Input: {input_path}")
print(f"Output: {output}")
if __name__ == "__main__":
main()