From 72250aa686c4be535f6d19038f9e0c13a0f2fb9e Mon Sep 17 00:00:00 2001 From: stanbot8 Date: Fri, 24 Jul 2026 08:21:03 -0700 Subject: [PATCH] fix: support raster labels above uint16 --- src/spatialdata/_core/operations/rasterize.py | 7 ++---- tests/core/operations/test_rasterize.py | 23 +++++++++++++++++++ 2 files changed, 25 insertions(+), 5 deletions(-) diff --git a/src/spatialdata/_core/operations/rasterize.py b/src/spatialdata/_core/operations/rasterize.py index d5b28c281..f2eb09539 100644 --- a/src/spatialdata/_core/operations/rasterize.py +++ b/src/spatialdata/_core/operations/rasterize.py @@ -30,7 +30,7 @@ get_axes_names, get_model, ) -from spatialdata.models._utils import get_spatial_axes +from spatialdata.models._utils import _get_uint_dtype, get_spatial_axes from spatialdata.transformations._utils import _get_scale, compute_coordinates from spatialdata.transformations.operations import get_transformation, remove_transformation, set_transformation from spatialdata.transformations.transformations import ( @@ -733,10 +733,7 @@ def rasterize_shapes_points( max_label = next(iter(reversed(label_index_to_category.keys()))) else: max_label = int(agg.max().values) - max_uint16 = np.iinfo(np.uint16).max - if max_label > max_uint16: - raise ValueError(f"Maximum label index is {max_label}. Values higher than {max_uint16} are not supported.") - agg = agg.astype(np.uint16) + agg = agg.astype(_get_uint_dtype(max_label)) return Labels2DModel.parse(agg, transformations=transformations) agg = agg.expand_dims(dim={"c": 1}).transpose("c", "y", "x") diff --git a/tests/core/operations/test_rasterize.py b/tests/core/operations/test_rasterize.py index 53d362d15..249ec94a2 100644 --- a/tests/core/operations/test_rasterize.py +++ b/tests/core/operations/test_rasterize.py @@ -446,6 +446,29 @@ def _rasterize(element: DaskDataFrame, **kwargs) -> SpatialImage: assert d[res[1, 2]] == 4 +def test_rasterize_points_selects_label_dtype(): + n_points = np.iinfo(np.uint16).max + 1 + points = PointsModel.parse( + dd.from_pandas( + pd.DataFrame({"x": np.arange(n_points), "y": np.zeros(n_points)}), + npartitions=1, + ) + ) + + result = rasterize( + points, + axes=("x", "y"), + min_coordinate=[0, 0], + max_coordinate=[n_points, 1], + target_coordinate_system="global", + target_width=n_points, + return_regions_as_labels=True, + ) + + assert result.dtype == np.uint32 + assert result.max().compute() == n_points + + def test_rasterize_spatialdata(full_sdata): sdata = full_sdata.subset( ["image2d", "image2d_multiscale", "labels2d", "labels2d_multiscale", "points_0", "circles"]