|
16 | 16 | _XGCM_AXES = Mapping[_XGCM_AXIS_DIRECTION, xgcm.Axis] |
17 | 17 |
|
18 | 18 |
|
19 | | -def get_tracer_dimensionality(axis: xgcm.Axis | None) -> int: |
| 19 | +def get_n_cell_edges_along_dim(axis: xgcm.Axis | None) -> int: |
20 | 20 | if axis is None: |
21 | 21 | return 1 |
22 | 22 | first_coord = list(axis.coords.items())[0] |
23 | | - pos, coord = first_coord |
| 23 | + _, coord_var = first_coord |
24 | 24 |
|
25 | | - pos_to_dim = { # TODO: These could do with being explicitly tested |
26 | | - "center": lambda x: x, |
27 | | - "left": lambda x: x, |
28 | | - "right": lambda x: x, |
29 | | - "inner": lambda x: x + 1, |
30 | | - "outer": lambda x: x - 1, |
31 | | - } |
32 | | - |
33 | | - n = axis._ds[coord].size |
34 | | - return pos_to_dim[pos](n) |
| 25 | + return axis._ds[coord_var].size |
35 | 26 |
|
36 | 27 |
|
37 | 28 | def get_time(axis: xgcm.Axis) -> npt.NDArray: |
@@ -121,22 +112,19 @@ def time(self): |
121 | 112 |
|
122 | 113 | @property |
123 | 114 | def xdim(self): |
124 | | - """Number of T (tracer) cells in the X direction.""" |
125 | | - return get_tracer_dimensionality(self.xgcm_grid.axes.get("X")) |
| 115 | + return get_n_cell_edges_along_dim(self.xgcm_grid.axes.get("X")) |
126 | 116 |
|
127 | 117 | @property |
128 | 118 | def ydim(self): |
129 | | - """Number of T (tracer) cells in the Y direction.""" |
130 | | - return get_tracer_dimensionality(self.xgcm_grid.axes.get("Y")) |
| 119 | + return get_n_cell_edges_along_dim(self.xgcm_grid.axes.get("Y")) |
131 | 120 |
|
132 | 121 | @property |
133 | 122 | def zdim(self): |
134 | | - """Number of T (tracer) cells in the Z direction.""" |
135 | | - return get_tracer_dimensionality(self.xgcm_grid.axes.get("Z")) |
| 123 | + return get_n_cell_edges_along_dim(self.xgcm_grid.axes.get("Z")) |
136 | 124 |
|
137 | 125 | @property |
138 | 126 | def tdim(self): |
139 | | - return get_tracer_dimensionality(self.xgcm_grid.axes.get("T")) |
| 127 | + return get_n_cell_edges_along_dim(self.xgcm_grid.axes.get("T")) |
140 | 128 |
|
141 | 129 | @property |
142 | 130 | def time_origin(self): |
|
0 commit comments