Skip to content

Commit 6d19769

Browse files
committed
Add unit test test_transpose_xfield_data_to_tzyx
1 parent 8b79856 commit 6d19769

3 files changed

Lines changed: 37 additions & 11 deletions

File tree

tests/utils.py

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,19 @@
11
"""General helper functions and utilies for test suite."""
22

3+
from __future__ import annotations
4+
35
from pathlib import Path
6+
from typing import TYPE_CHECKING
47

58
import numpy as np
69
import xarray as xr
710

811
import parcels
912
from parcels import FieldSet
13+
from parcels.xgrid import _FIELD_DATA_ORDERING, get_axis_from_dim_name
14+
15+
if TYPE_CHECKING:
16+
from parcels.xgrid import XGrid
1017

1118
PROJECT_ROOT = Path(__file__).resolve().parents[1]
1219
TEST_ROOT = PROJECT_ROOT / "tests"
@@ -116,3 +123,13 @@ def create_fieldset_zeros_simple(xdim=40, ydim=100, withtime=False):
116123

117124
def assert_empty_folder(path: Path):
118125
assert [p.name for p in path.iterdir()] == []
126+
127+
128+
def assert_valid_field_data(data: xr.DataArray, grid: XGrid):
129+
assert len(data.shape) == 4, f"Field data should have 4 dimensions (time, depth, lat, lon), got {len(data.shape)}"
130+
131+
for ax_expected, dim in zip(_FIELD_DATA_ORDERING, data.dims, strict=True):
132+
ax_actual = get_axis_from_dim_name(grid.xgcm_grid.axes, dim)
133+
if ax_actual is None:
134+
continue # None is ok
135+
assert ax_actual == ax_expected, f"Expected axis {ax_expected} for dimension '{dim}', got {ax_actual}"

tests/v4/test_fieldset.py

Lines changed: 3 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,8 @@
1212
from parcels._datasets.structured.generic import datasets as datasets_structured
1313
from parcels.field import Field, VectorField
1414
from parcels.fieldset import CalendarError, FieldSet, _datetime_to_msg
15-
from parcels.xgrid import _FIELD_DATA_ORDERING, XGrid, get_axis_from_dim_name
15+
from parcels.xgrid import XGrid
16+
from tests import utils
1617

1718
ds = datasets_structured["ds_2d_left"]
1819

@@ -95,15 +96,7 @@ def test_fieldset_from_structured_generic_datasets(ds):
9596

9697
assert len(fieldset.fields) == len(ds.data_vars)
9798
for field in fieldset.fields.values():
98-
assert (
99-
len(field.data.shape) == 4
100-
), f"Field data should have 4 dimensions (time, depth, lat, lon), got {len(field.data.shape)}"
101-
102-
for ax_expected, dim in zip(_FIELD_DATA_ORDERING, field.data.dims, strict=True):
103-
ax_actual = get_axis_from_dim_name(field.grid.xgcm_grid.axes, dim)
104-
if ax_actual is None:
105-
continue # None is ok
106-
assert ax_actual == ax_expected, f"Expected axis {ax_expected} for dimension '{dim}', got {ax_actual}"
99+
utils.assert_valid_field_data(field.data, field.grid)
107100

108101
assert len(fieldset.gridset) == 1
109102

tests/v4/test_xgrid.py

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,4 @@
1+
import itertools
12
from collections import namedtuple
23

34
import numpy as np
@@ -6,7 +7,8 @@
67
from numpy.testing import assert_allclose
78

89
from parcels._datasets.structured.generic import X, Y, Z, datasets
9-
from parcels.xgrid import XGrid, _drop_field_data, _search_1d_array
10+
from parcels.xgrid import XGrid, _drop_field_data, _search_1d_array, _transpose_xfield_data_to_tzyx
11+
from tests import utils
1012

1113
GridTestCase = namedtuple("GridTestCase", ["ds", "attr", "expected"])
1214

@@ -49,6 +51,20 @@ def test_xgrid_axes(ds):
4951
assert grid.axes == ["Z", "Y", "X"]
5052

5153

54+
@pytest.mark.parametrize("ds", [datasets["ds_2d_left"]])
55+
def test_transpose_xfield_data_to_tzyx(ds):
56+
da = ds["data_g"]
57+
grid = XGrid.from_dataset(ds)
58+
59+
all_combinations = (itertools.combinations(da.dims, n) for n in range(len(da.dims)))
60+
all_combinations = itertools.chain(*all_combinations)
61+
for subset_dims in all_combinations:
62+
isel = {dim: 0 for dim in subset_dims}
63+
da_subset = da.isel(isel, drop=True)
64+
da_test = _transpose_xfield_data_to_tzyx(da_subset, grid.xgcm_grid)
65+
utils.assert_valid_field_data(da_test, grid)
66+
67+
5268
@pytest.mark.parametrize("ds", [datasets["ds_2d_left"]])
5369
def test_xgrid_get_axis_dim(ds):
5470
grid = XGrid.from_dataset(ds)

0 commit comments

Comments
 (0)