Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 28 additions & 18 deletions packages/evaluate/src/weathergen/evaluate/plotting/plot_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@
from collections.abc import Iterable, Sequence

import numpy as np
import xarray as xr

_logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -99,17 +100,19 @@ def plot_metric_region(
if ch not in np.atleast_1d(data.channel.values) or data.isnull().all():
continue

data, time_dim = _assign_time_coord(data)

selected_data.append(data.sel(channel=ch))
labels.append(runs[run_id].get("label", run_id))
run_ids.append(run_id)

if selected_data:
_logger.info(f"Creating plot for {metric} - {region} - {stream} - {ch}.")

name = create_filename(
prefix=[metric, region], middle=sorted(set(run_ids)), suffix=[stream, ch]
)

selected_data, time_dim = _assign_time_coord(selected_data)

plotter.plot(
selected_data,
labels,
Expand All @@ -120,12 +123,12 @@ def plot_metric_region(
)


def _assign_time_coord(data: object) -> object:
def _assign_time_coord(selected_data: list[xr.DataArray]) -> tuple[xr.DataArray, str]:
"""Ensure that lead_time coordinate exists in the data array.

Parameters
----------
data : xarray.DataArray
selected_data : list[xarray.DataArray]
The data array to check.

Returns
Expand All @@ -136,23 +139,30 @@ def _assign_time_coord(data: object) -> object:
time_dim : str
The name of the time dimension used for x-axis.
"""
if "forecast_step" not in data.dims and "forecast_step" not in data.coords:
raise ValueError("forecast_step coordinate not found in data dimensions or coordinates.")

time_dim = "forecast_step"

if "lead_time" in data.coords and data["forecast_step"].size == data["lead_time"].size:
data = data.swap_dims({"forecast_step": "lead_time"})
for data in selected_data:
if "forecast_step" not in data.dims and "forecast_step" not in data.coords:
raise ValueError(
"forecast_step coordinate not found in data dimensions or coordinates."
)

# Prefer lead_time as x_dim if present in dimensions
if "lead_time" in data.dims:
time_dim = "lead_time"
else:
_logger.warning(
"lead_time coordinate not found or mismatched size; using forecast_step as x-axis."
)
if "lead_time" not in data.coords and "lead_time" not in data.dims:
_logger.warning(
"lead_time coordinate not found for all plotted data; "
"using forecast_step as x-axis."
)
return selected_data, time_dim

# Swap forecast_step with lead_time if all available run_ids have lead_time coord
time_dim = "lead_time"

return data, time_dim
for i, data in enumerate(selected_data):
if data.coords["lead_time"].shape == data.coords["forecast_step"].shape:
selected_data[i] = data.swap_dims({"forecast_step": "lead_time"})

return selected_data, time_dim


def ratio_plot_metric_region(
Expand Down Expand Up @@ -251,8 +261,6 @@ def heat_maps_metric_region(
if data.isnull().all():
continue

data, time_dim = _assign_time_coord(data)

selected_data.append(data)
label = runs[run_id].get("label", run_id)
if label != run_id:
Expand All @@ -265,6 +273,8 @@ def heat_maps_metric_region(
name = create_filename(
prefix=[metric, region], middle=sorted(set(run_ids)), suffix=[stream]
)
selected_data, time_dim = _assign_time_coord(selected_data)

plotter.heat_map(
selected_data,
labels,
Expand Down
17 changes: 9 additions & 8 deletions packages/evaluate/src/weathergen/evaluate/utils/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -214,14 +214,15 @@ def calc_scores_per_stream(
reader, map_dir, stream, region, score_data, metrics, fstep
)

lead_time_values = np.array(
[lead_time_map[f].astype(int) for f in metric_stream.forecast_step.values]
).squeeze()

if lead_time_values.shape == metric_stream.forecast_step.shape:
metric_stream = metric_stream.assign_coords(
lead_time=("forecast_step", lead_time_values)
)
if all(lead_time_map[f] is not None for f in lead_time_map):
lead_time_values = np.array(
[lead_time_map[f].astype(int) for f in metric_stream.forecast_step.values]
).squeeze()

if lead_time_values.shape == metric_stream.forecast_step.shape:
metric_stream = metric_stream.assign_coords(
lead_time=("forecast_step", lead_time_values)
)

_logger.info(f"Scores for run {reader.run_id} - {stream} calculated successfully.")

Expand Down