Skip to content

Commit f6228ff

Browse files
committed
Add get_axis_dim and axes to BaseGrid, XGrid, and UxGrid
Also added some simple tests in test_uxgrid.py
1 parent ce49751 commit f6228ff

5 files changed

Lines changed: 102 additions & 21 deletions

File tree

parcels/basegrid.py

Lines changed: 35 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -135,3 +135,38 @@ def unravel_index(self, ei: int) -> dict[str, int]:
135135
Raised if the method is not implemented for the current grid type.
136136
"""
137137
...
138+
139+
@property
140+
@abstractmethod
141+
def axes(self) -> list[str]:
142+
"""
143+
Return a list of axis names that are part of this grid.
144+
145+
Returns
146+
-------
147+
list[str]
148+
List of axis names, e.g. ["Z", "Y", "X"] for a 3D structured grid or ["Z", "FACE"] for an unstructured grid.
149+
"""
150+
...
151+
152+
@abstractmethod
153+
def get_axis_dim(self, axis: str) -> int:
154+
"""
155+
Return the dimensionality (number of cells/edges) along a specific axis.
156+
157+
Parameters
158+
----------
159+
axis : str
160+
The name of the axis to get the dimensionality for. Must be one of the values returned by self.axes.
161+
162+
Returns
163+
-------
164+
int
165+
The number of cells/edges along the specified axis.
166+
167+
Raises
168+
------
169+
ValueError
170+
If the specified axis is not part of this grid.
171+
"""
172+
...

parcels/uxgrid.py

Lines changed: 13 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -54,6 +54,19 @@ def depth(self):
5454
return np.zeros(1)
5555
return self.z.values
5656

57+
@property
58+
def axes(self) -> list[_UXGRID_AXES]:
59+
return ["Z", "FACE"]
60+
61+
def get_axis_dim(self, axis: _UXGRID_AXES) -> int:
62+
if axis not in self.axes:
63+
raise ValueError(f"Axis {axis!r} is not part of this grid. Available axes: {self.axes}")
64+
65+
if axis == "Z":
66+
return len(self.z.values)
67+
elif axis == "FACE":
68+
return self.uxgrid.n_face
69+
5770
def search(self, z, y, x, ei=None):
5871
tol = 1e-10
5972

parcels/xgrid.py

Lines changed: 18 additions & 13 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
from collections.abc import Hashable, Mapping
2+
from functools import cached_property
23
from typing import Literal, cast
34

45
import numpy as np
@@ -112,21 +113,23 @@ def _datetimes(self):
112113
def time(self):
113114
return self._datetimes.astype(np.float64) / 1e9
114115

115-
@property
116-
def xdim(self):
117-
return get_cell_edge_count_along_dim(self.xgcm_grid.axes.get("X"))
116+
@cached_property
117+
def xdim(self) -> int:
118+
return self.get_axis_dim("X")
118119

119-
@property
120-
def ydim(self):
121-
return get_cell_edge_count_along_dim(self.xgcm_grid.axes.get("Y"))
120+
@cached_property
121+
def ydim(self) -> int:
122+
return self.get_axis_dim("Y")
122123

123-
@property
124-
def zdim(self):
125-
return get_cell_edge_count_along_dim(self.xgcm_grid.axes.get("Z"))
124+
@cached_property
125+
def zdim(self) -> int:
126+
return self.get_axis_dim("Z")
126127

127-
@property
128-
def tdim(self):
129-
return get_cell_edge_count_along_dim(self.xgcm_grid.axes.get("T"))
128+
def get_axis_dim(self, axis: _XGRID_AXES) -> int:
129+
if axis not in self.axes:
130+
raise ValueError(f"Axis {axis!r} is not part of this grid. Available axes: {self.axes}")
131+
132+
return get_cell_edge_count_along_dim(self.xgcm_grid.axes.get(axis))
130133

131134
@property
132135
def _z4d(self) -> Literal[0, 1]:
@@ -188,7 +191,9 @@ def ravel_index(self, axis_indices: dict[_XGRID_AXES, int]) -> int:
188191
xi = axis_indices.get("X", 0)
189192
yi = axis_indices.get("Y", 0)
190193
zi = axis_indices.get("Z", 0)
191-
return xi + self.xdim * yi + self.xdim * self.ydim * zi
194+
xdim = self.get_axis_dim("X")
195+
ydim = self.get_axis_dim("Y")
196+
return xi + xdim * yi + xdim * ydim * zi
192197

193198
def unravel_index(self, ei) -> dict[_XGRID_AXES, int]:
194199
zi = ei // (self.xdim * self.ydim)

tests/v4/test_uxgrid.py

Lines changed: 23 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,23 @@
1+
import pytest
2+
3+
from parcels._datasets.unstructured.generic import datasets as uxdatasets
4+
from parcels.uxgrid import UxGrid
5+
6+
7+
@pytest.mark.parametrize("uxds", [pytest.param(uxds, id=key) for key, uxds in uxdatasets.items()])
8+
def test_uxgrid_init_on_generic_datasets(uxds):
9+
UxGrid(uxds.uxgrid, z=uxds.coords["nz"])
10+
11+
12+
@pytest.mark.parametrize("uxds", [uxdatasets["stommel_gyre_delaunay"]])
13+
def test_uxgrid_axes(uxds):
14+
grid = UxGrid(uxds.uxgrid, z=uxds.coords["nz"])
15+
assert grid.axes == ["Z", "FACE"]
16+
17+
18+
@pytest.mark.parametrize("uxds", [uxdatasets["stommel_gyre_delaunay"]])
19+
def test_xgrid_get_axis_dim(uxds):
20+
grid = UxGrid(uxds.uxgrid, z=uxds.coords["nz"])
21+
22+
assert grid.get_axis_dim("FACE") == 721
23+
assert grid.get_axis_dim("Z") == 2

tests/v4/test_xgrid.py

Lines changed: 13 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,7 @@
66
from numpy.testing import assert_allclose
77

88
from parcels import xgcm
9-
from parcels._datasets.structured.generic import T, X, Y, Z, datasets
9+
from parcels._datasets.structured.generic import X, Y, Z, datasets
1010
from parcels.xgrid import XGrid, _search_1d_array
1111

1212
GridTestCase = namedtuple("GridTestCase", ["Grid", "attr", "expected"])
@@ -16,10 +16,6 @@
1616
GridTestCase(datasets["ds_2d_left"], "lat", datasets["ds_2d_left"].YG.values),
1717
GridTestCase(datasets["ds_2d_left"], "depth", datasets["ds_2d_left"].ZG.values),
1818
GridTestCase(datasets["ds_2d_left"], "time", datasets["ds_2d_left"].time.values.astype(np.float64) / 1e9),
19-
GridTestCase(datasets["ds_2d_left"], "xdim", X),
20-
GridTestCase(datasets["ds_2d_left"], "ydim", Y),
21-
GridTestCase(datasets["ds_2d_left"], "zdim", Z),
22-
GridTestCase(datasets["ds_2d_left"], "tdim", T),
2319
]
2420

2521

@@ -45,9 +41,18 @@ def test_xgrid_init_on_generic_datasets(ds):
4541
XGrid(xgcm.Grid(ds, periodic=False))
4642

4743

48-
def test_xgrid_axes():
49-
# Tests that the xgrid.axes property correctly identifies the axes and ordering
50-
...
44+
@pytest.mark.parametrize("ds", [datasets["ds_2d_left"]])
45+
def test_xgrid_axes(ds):
46+
grid = XGrid(xgcm.Grid(ds, periodic=False))
47+
assert grid.axes == ["Z", "Y", "X"]
48+
49+
50+
@pytest.mark.parametrize("ds", [datasets["ds_2d_left"]])
51+
def test_xgrid_get_axis_dim(ds):
52+
grid = XGrid(xgcm.Grid(ds, periodic=False))
53+
assert grid.get_axis_dim("Z") == Z
54+
assert grid.get_axis_dim("Y") == Y
55+
assert grid.get_axis_dim("X") == X
5156

5257

5358
def test_invalid_xgrid_field_array():

0 commit comments

Comments
 (0)