Skip to content

Commit 8ced3da

Browse files
committed
done
1 parent ed4a22b commit 8ced3da

3 files changed

Lines changed: 27 additions & 3 deletions

File tree

CHANGES.md

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -2,7 +2,9 @@
22

33
- Fixed an issue in `rectify_dataset` when processing datasets with decreasing
44
x-coordinates.
5-
5+
- Boolean (`bool`) data arrays now default to nearest-neighbor interpolation and
6+
center for spatial aggregation, matching the behavior of integer arrays to
7+
ensure valid results.
68

79
## Changes in 0.3.0
810

tests/test_utils.py

Lines changed: 14 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -117,6 +117,9 @@ def test_n_get_grid_mapping_name(self):
117117

118118
def test_get_interp_method(self):
119119
int_var = xr.DataArray(np.array([1, 2, 3], dtype=np.int32), dims=["x"])
120+
bool_var = xr.DataArray(
121+
np.array([True, False, True], dtype=np.bool_), dims=["x"]
122+
)
120123
float_var = xr.DataArray(
121124
np.array([1.0, 2.0, 3.0], dtype=np.float32), dims=["x"]
122125
)
@@ -125,6 +128,10 @@ def test_get_interp_method(self):
125128
result = _get_spatial_interp_method(None, "var", int_var)
126129
self.assertEqual(result, 0)
127130

131+
# bool type data array
132+
result = _get_spatial_interp_method(None, "var", bool_var)
133+
self.assertEqual(result, 0)
134+
128135
# float type data array
129136
result = _get_spatial_interp_method(None, "var", float_var)
130137
self.assertEqual(result, 1)
@@ -182,6 +189,9 @@ def test_prep_interp_methods_downscale(self):
182189

183190
def test_get_agg_method(self):
184191
int_var = xr.DataArray(np.array([1, 2, 3], dtype=np.int32), dims=["x"])
192+
bool_var = xr.DataArray(
193+
np.array([True, False, True], dtype=np.bool_), dims=["x"]
194+
)
185195
float_var = xr.DataArray(
186196
np.array([1.0, 2.0, 3.0], dtype=np.float32), dims=["x"]
187197
)
@@ -190,6 +200,10 @@ def test_get_agg_method(self):
190200
result = _get_spatial_agg_method(None, "var", int_var)
191201
self.assertEqual(result, AGG_METHODS["center"])
192202

203+
# bool type data array, default
204+
result = _get_spatial_agg_method(None, "var", bool_var)
205+
self.assertEqual(result, AGG_METHODS["center"])
206+
193207
# float type data array, default
194208
result = _get_spatial_agg_method(None, "var", float_var)
195209
self.assertEqual(result, AGG_METHODS["mean"])

xcube_resampling/utils.py

Lines changed: 10 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -379,7 +379,11 @@ def _get_spatial_interp_method(
379379
var: xr.DataArray,
380380
) -> SpatialInterpMethod:
381381
def assign_defaults(data_type: np.dtype) -> SpatialInterpMethod:
382-
return 0 if np.issubdtype(data_type, np.integer) else 1
382+
if np.issubdtype(data_type, np.bool_):
383+
return 0
384+
if np.issubdtype(data_type, np.integer):
385+
return 0
386+
return 1
383387

384388
if isinstance(interp_methods, Mapping):
385389
interp_method = interp_methods.get(str(key), interp_methods.get(var.dtype))
@@ -450,7 +454,11 @@ def _get_spatial_agg_method(
450454
var: xr.DataArray,
451455
) -> Callable:
452456
def assign_defaults(data_type: np.dtype) -> SpatialAggMethod:
453-
return "center" if np.issubdtype(data_type, np.integer) else "mean"
457+
if np.issubdtype(data_type, np.bool_):
458+
return "center"
459+
if np.issubdtype(data_type, np.integer):
460+
return "center"
461+
return "mean"
454462

455463
if isinstance(agg_methods, Mapping):
456464
agg_method = agg_methods.get(str(key), agg_methods.get(var.dtype))

0 commit comments

Comments
 (0)