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
4 changes: 3 additions & 1 deletion src/earthkit/data/indexing/tensor.py
Original file line number Diff line number Diff line change
Expand Up @@ -503,7 +503,9 @@ def make_valid_datetime(self, dims_map, dtype="datetime64[ns]"):
dims_step = [keys[d] for d in dims]
# use same dim order as in user_dims
dims = [d for d in self.user_dims if d in dims_step]
assert len(dims) == len(dims_step), f"Duplicate dims in {dims}"
if len(dims) != len(dims_step):
continue
assert len(dims) == len(dims_step), f"{dims=} {dims_step=}"
other_dims = [d for d in self.user_dims if d not in dims]

if other_dims:
Expand Down
28 changes: 14 additions & 14 deletions src/earthkit/data/utils/xarray/check.py
Original file line number Diff line number Diff line change
Expand Up @@ -48,17 +48,20 @@ def __init__(self, tensor):
self.tensor = tensor

def first_diff(self, coord_keys):
for i, f in enumerate(self.tensor.source):
t_coords = self.tensor._index_to_coords_value(i, self.tensor)
f_coords = f.metadata(coord_keys)
print(f"Checking field[{i}] with coords {t_coords} vs {f_coords}")
diff = ListDiff.diff(t_coords, f_coords)
if not diff.same:
name = ""
if diff.diff_index != -1:
name = coord_keys[diff.diff_index]

return i, f, t_coords, f_coords, name, diff
if coord_keys:
for i, f in enumerate(self.tensor.source):
t_coords = self.tensor._index_to_coords_value(i, self.tensor)
f_coords = f.metadata(coord_keys)
try:
diff = ListDiff.diff(t_coords, f_coords)
if not diff.same:
name = ""
if diff.diff_index != -1:
name = coord_keys[diff.diff_index]

return i, f, t_coords, f_coords, name, diff
except Exception:
pass

def neighbour_field(self, field_num, index):
f_other = None
Expand Down Expand Up @@ -104,9 +107,6 @@ def check(self, details=False):
f"Dimensions: \n {dims}"
)

print(text_num)
print(coord_keys)

if not details:
raise ValueError(text_num)
else:
Expand Down
60 changes: 46 additions & 14 deletions src/earthkit/data/utils/xarray/grib.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,15 +16,18 @@
from earthkit.data.utils.dates import step_to_grib
from earthkit.data.utils.dates import time_to_grib
from earthkit.data.utils.dates import to_datetime
from earthkit.data.utils.dates import to_timedelta

LOG = logging.getLogger(__name__)


def update_metadata(metadata, compulsory):
def update_metadata(metadata, compulsory, step_len=0):
if "valid_time" in metadata:
dt = to_datetime(metadata.pop("valid_time"))
metadata["date"] = dt.date()
metadata["time"] = dt.time()
metadata["stepRange"] = step_to_grib(metadata.pop("stepRange", 0))
metadata.pop("step", None)

if "forecast_reference_time" in metadata:
date, time = datetime_to_grib(to_datetime(metadata["forecast_reference_time"]))
Expand All @@ -43,7 +46,19 @@ def update_metadata(metadata, compulsory):
metadata["time"] = date.hour * 100 + date.minute

if "step" in metadata:
metadata["step"] = step_to_grib(metadata["step"])
if step_len is None:
metadata["step"] = step_to_grib(metadata["step"])
elif step_len.total_seconds() == 0:
step = step_to_grib(metadata["step"])
metadata["stepRange"] = step
metadata.pop("step", None)
elif step_len.total_seconds() > 0:
end = metadata["step"]
start = to_timedelta(end) - step_len
start = step_to_grib(start)
end = step_to_grib(end)
metadata["stepRange"] = f"{start}-{end}"
metadata.pop("step", None)

if "stream" not in metadata:
if "number" in metadata:
Expand Down Expand Up @@ -104,7 +119,28 @@ def data_array_to_fields(da, metadata=None):
else:
coords[k] = coords[k].values

# print(f"{coords=}")
# extract metadata template from dataarray
if hasattr(da, "earthkit"):
template_metadata = da.earthkit.metadata
else:
raise ValueError("Earthkit attribute not found in DataArray. Required for conversion to FieldList!")

# special treatment to set step related GRIB keys
compulsory_metadata = {}
step_len = datetime.timedelta(hours=0)
if "valid_time" in dims:
# when valid_time is a dimension we enforce the step to be 0
compulsory_metadata["stepRange"] = 0
else:
try:
step_range = template_metadata.get("stepRange", None)
if isinstance(step_range, str) and "-" in step_range:
step_len = to_timedelta(template_metadata.get("endStep", 0)) - to_timedelta(
template_metadata.get("startStep", 0)
)
except TypeError as e:
print(f"Error calculating step length: {e}")
step_len = None

for values in product(*[coords[dim] for dim in dims]):

Expand All @@ -124,16 +160,12 @@ def data_array_to_fields(da, metadata=None):
grib_metadata.update(dict(zip(components[k], grib_metadata[k][1:])))
# print(f"-> {grib_metadata=}")
del grib_metadata[k]
update_metadata(grib_metadata, [])

# extract metadata from object
if metadata is None:
if hasattr(da, "earthkit"):
metadata = da.earthkit.metadata
else:
raise ValueError(
"Earthkit attribute not found in DataArray. Required for conversion to FieldList!"
)

metadata = metadata.override(grib_metadata)
for k in compulsory_metadata:
if k not in grib_metadata:
grib_metadata[k] = compulsory_metadata[k]

update_metadata(grib_metadata, [], step_len=step_len)

metadata = template_metadata.override(grib_metadata)
yield ArrayField(xa_field.values, metadata)
20 changes: 20 additions & 0 deletions tests/xr_engine/test_xr_time.py
Original file line number Diff line number Diff line change
Expand Up @@ -574,3 +574,23 @@ def test_xr_time_step_range_2(kwargs, dims, step_units):
assert (
ds[step_units[0]].attrs["units"] == step_units[1]
), f"step units mismatch {ds[step_units[0]].attrs['units']} != {step_units[1]}"


@pytest.mark.cache
def test_xr_time_forecast_per_month():
ds_ek = from_source("url", earthkit_remote_test_data_file("xr_engine/date/2_months_6_hourly.grib"))

ds = ds_ek.to_xarray(time_dim_mode="valid_time")

ref = []
start = np.datetime64("1979-01-01T06:00:00", "ns")
end = np.datetime64("1979-03-01T00:00:00", "ns")
while start <= end:
ref.append(np.datetime64(start))
start += np.timedelta64(6, "h")

dims = {
"valid_time": ref,
}

compare_dims(ds, dims, order_ref_var="avg_dis")
32 changes: 32 additions & 0 deletions tests/xr_engine/test_xr_write.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from earthkit.data import to_target
from earthkit.data.core.temporary import temp_file
from earthkit.data.testing import earthkit_remote_test_data_file
from earthkit.data.utils.dates import datetime_to_grib


@pytest.mark.cache
Expand Down Expand Up @@ -457,3 +458,34 @@ def test_xr_write_to_netcdf_file_dataset(method, kwargs):

assert np.allclose(ref_t_vals + 1.0, r["t"].isel(step=1, level=0).to_numpy().flatten())
assert np.allclose(ref_r_vals + 1.0, r["r"].isel(step=1, level=0).to_numpy().flatten())


@pytest.mark.cache
def test_xr_write_forecast_per_month():
ds_ek = from_source("url", earthkit_remote_test_data_file("xr_engine/date/2_months_6_hourly.grib"))

ds = ds_ek.to_xarray(time_dim_mode="valid_time")
r = ds.earthkit.to_fieldlist()
assert len(r) == 236

ref = []
start = np.datetime64("1979-01-01T06:00:00", "ns")
end = np.datetime64("1979-03-01T00:00:00", "ns")
while start <= end:
base_date, base_time = datetime_to_grib(start)
ref.append([base_date, base_time, 0, "0", 0, 0, base_date, base_time])
start += np.timedelta64(6, "h")

keys = [
"dataDate",
"dataTime",
"step",
"stepRange",
"startStep",
"endStep",
"validityDate",
"validityTime",
]

for f, f_ref in zip(r, ref):
assert f.metadata(keys) == f_ref, f"Expected: {f_ref}\nGot: {f.metadata(keys)}"
Loading