|
10 | 10 | from parcels.basegrid import BaseGrid |
11 | 11 | from parcels.tools.converters import TimeConverter |
12 | 12 |
|
13 | | -_AXIS_DIRECTION = Literal["X", "Y", "Z", "T"] |
14 | | -_AXIS_POSITION = Literal["center", "left", "right", "inner", "outer"] |
15 | | -_XGCM_AXES = Mapping[_AXIS_DIRECTION, xgcm.Axis] |
| 13 | +_XGCM_AXIS_DIRECTION = Literal["X", "Y", "Z", "T"] |
| 14 | +_XGCM_AXIS_POSITION = Literal["center", "left", "right", "inner", "outer"] |
| 15 | +_AXIS_DIRECTION = Literal["X", "Y", "Z"] |
| 16 | +_XGCM_AXES = Mapping[_XGCM_AXIS_DIRECTION, xgcm.Axis] |
16 | 17 |
|
17 | 18 |
|
18 | 19 | def get_tracer_dimensionality(axis: xgcm.Axis | None) -> int: |
@@ -180,85 +181,57 @@ def _gtype(self): |
180 | 181 | else: |
181 | 182 | return GridType.CurvilinearSGrid |
182 | 183 |
|
183 | | - def search(self, z, y, x, ei=None, search2D=False): |
| 184 | + def search(self, z, y, x, ei=None): |
184 | 185 | ds = self.xgcm_grid._ds |
185 | 186 |
|
186 | | - if search2D: |
187 | | - zi = 0 |
188 | | - else: |
189 | | - zi, _ = _search_1d_array(ds.depth.values, z) |
| 187 | + zi, zeta = _search_1d_array(ds.depth.values, z) |
190 | 188 |
|
191 | 189 | if ds.lon.ndim == 1: |
192 | 190 | yi, eta = _search_1d_array(ds.lat.values, y) |
193 | 191 | xi, xsi = _search_1d_array(ds.lon.values, x) |
194 | | - return np.array([eta, xsi, 1 - eta, 1 - xsi]), self.ravel_index(zi, yi, xi) |
| 192 | + return {"X": (xi, xsi), "Y": (yi, eta), "Z": (zi, zeta)} |
195 | 193 |
|
196 | 194 | yi, xi = None, None |
197 | 195 | if ei is not None: |
198 | | - _, yi, xi = self.unravel_index(ei) |
| 196 | + axis_indices = self.unravel_index(ei) |
| 197 | + xi = axis_indices.get("X") |
| 198 | + yi = axis_indices.get("Y") |
199 | 199 |
|
200 | 200 | if ds.lon.ndim == 2: |
201 | 201 | eta, xsi, yi, xi = _search_indices_curvilinear_2d(self, y, x, yi, xi) |
202 | 202 |
|
203 | | - return np.array([eta, xsi, 1 - eta, 1 - xsi]), self.ravel_index(zi, yi, xi) |
| 203 | + return {"X": (xi, xsi), "Y": (yi, eta), "Z": (zi, zeta)} |
204 | 204 |
|
205 | 205 | raise NotImplementedError("Searching in >2D lon/lat arrays is not implemented yet.") |
206 | 206 |
|
207 | | - def ravel_index(self, zi, yi, xi): |
208 | | - """ |
209 | | - Converts a z, y, and x index into a single encoded index. |
210 | | -
|
211 | | - Parameters |
212 | | - ---------- |
213 | | - zi : int |
214 | | - Vertical index. |
215 | | - yi : int |
216 | | - Latitude index. |
217 | | - xi : int |
218 | | - Longitude index. |
219 | | -
|
220 | | - Returns |
221 | | - ------- |
222 | | - int |
223 | | - Encoded index. |
224 | | - """ |
| 207 | + def ravel_index(self, axis_indices: dict[_AXIS_DIRECTION, int]) -> int: |
| 208 | + xi = axis_indices.get("X", 0) |
| 209 | + yi = axis_indices.get("Y", 0) |
| 210 | + zi = axis_indices.get("Z", 0) |
225 | 211 | return xi + self.xdim * yi + self.xdim * self.ydim * zi |
226 | 212 |
|
227 | | - def unravel_index(self, ei): |
228 | | - """ |
229 | | - Converts a single encoded index back into a Z, Y, and X indices. |
230 | | -
|
231 | | - Parameters |
232 | | - ---------- |
233 | | - ei : int |
234 | | - Encoded index to be unraveled. |
235 | | -
|
236 | | - Returns |
237 | | - ------- |
238 | | - zi : int |
239 | | - Vertical index. |
240 | | - yi : int |
241 | | - Latitude index. |
242 | | - xi : int |
243 | | - Longitude index. |
244 | | - """ |
| 213 | + def unravel_index(self, ei) -> dict[_AXIS_DIRECTION, int]: |
245 | 214 | zi = ei // (self.xdim * self.ydim) |
246 | 215 | ei = ei % (self.xdim * self.ydim) |
247 | 216 |
|
248 | 217 | yi = ei // self.xdim |
249 | 218 | xi = ei % self.xdim |
250 | | - return zi, yi, xi |
| 219 | + return { |
| 220 | + "X": xi, |
| 221 | + "Y": yi, |
| 222 | + "Z": zi, |
| 223 | + } |
251 | 224 |
|
252 | 225 |
|
253 | | -def get_axis_from_dim_name(axes: _XGCM_AXES, dim: str) -> _AXIS_DIRECTION | None: |
| 226 | +def get_axis_from_dim_name(axes: _XGCM_AXES, dim: str) -> _XGCM_AXIS_DIRECTION | None: |
254 | 227 | """For a given dimension name in a grid, returns the direction axis it is on.""" |
255 | 228 | for axis_name, axis in axes.items(): |
256 | 229 | if dim in axis.coords.values(): |
257 | 230 | return axis_name |
258 | 231 | return None |
259 | 232 |
|
260 | 233 |
|
261 | | -def get_position_from_dim_name(axes: _XGCM_AXES, dim: str) -> _AXIS_POSITION | None: |
| 234 | +def get_position_from_dim_name(axes: _XGCM_AXES, dim: str) -> _XGCM_AXIS_POSITION | None: |
262 | 235 | """For a given dimension, returns the position of the variable in the grid.""" |
263 | 236 | for axis in axes.values(): |
264 | 237 | var_to_position = {var: position for position, var in axis.coords.items()} |
@@ -287,7 +260,7 @@ def assert_valid_field_array(da: xr.DataArray, axes: _XGCM_AXES): |
287 | 260 | assert_all_dimensions_correspond_with_axis(da, axes) |
288 | 261 |
|
289 | 262 | dim_to_axis = {dim: get_axis_from_dim_name(axes, dim) for dim in da.dims} |
290 | | - dim_to_axis = cast(dict[Hashable, _AXIS_DIRECTION], dim_to_axis) |
| 263 | + dim_to_axis = cast(dict[Hashable, _XGCM_AXIS_DIRECTION], dim_to_axis) |
291 | 264 |
|
292 | 265 | # Assert all dimensions are present |
293 | 266 | if set(dim_to_axis.values()) != {"T", "Z", "Y", "X"}: |
|
0 commit comments