diff --git a/src/spatialdata/_core/operations/transform.py b/src/spatialdata/_core/operations/transform.py index 410e92fc3..65fe9859a 100644 --- a/src/spatialdata/_core/operations/transform.py +++ b/src/spatialdata/_core/operations/transform.py @@ -360,10 +360,10 @@ def _( from spatialdata.transformations import get_transformation, set_transformation from spatialdata.transformations.transformations import Identity, Sequence - # labels need to be preserved after the resizing of the image + # labels hold categorical ids, so they must be resampled with nearest neighbour (order=0): interpolating would + # create ids that are not present in the input, as for single-scale labels in the DataArray overload above if schema in (Labels2DModel, Labels3DModel): - # TODO: this should work, test better - kwargs = {"prefilter": False} + kwargs = {"prefilter": False, "order": 0} channel_names = None elif schema in (Image2DModel, Image3DModel): kwargs = {} diff --git a/tests/core/operations/test_transform.py b/tests/core/operations/test_transform.py index ef307ac7b..bf980fe26 100644 --- a/tests/core/operations/test_transform.py +++ b/tests/core/operations/test_transform.py @@ -16,7 +16,7 @@ from spatialdata._core.data_extent import are_extents_equal, get_extent from spatialdata._core.spatialdata import SpatialData from spatialdata._utils import disable_dask_tune_optimization, unpad_raster -from spatialdata.models import Image2DModel, PointsModel, ShapesModel, get_axes_names +from spatialdata.models import Image2DModel, Labels2DModel, Labels3DModel, PointsModel, ShapesModel, get_axes_names from spatialdata.transformations.operations import ( align_elements_using_landmarks, get_transformation, @@ -234,6 +234,58 @@ def test_transform_shapes(shapes: SpatialData): assert geom_almost_equals(p0["geometry"], p1["geometry"]) +def _label_ids_per_scale(labels: DataArray | DataTree) -> list[set[int]]: + levels = [labels] if isinstance(labels, DataArray) else [next(iter(scale.values())) for scale in labels.values()] + return [set(np.unique(np.asarray(level.data)).tolist()) for level in levels] + + +def _rotation_xy(degrees: float) -> Affine: + theta = np.deg2rad(degrees) + matrix = np.array([[np.cos(theta), -np.sin(theta), 0], [np.sin(theta), np.cos(theta), 0], [0, 0, 1]]) + return Affine(matrix, input_axes=("x", "y"), output_axes=("x", "y")) + + +@pytest.mark.parametrize("multiscale", [False, True]) +@pytest.mark.parametrize("via_spatialdata", [False, True]) +@pytest.mark.parametrize( + "transformation", + [_rotation_xy(30), Scale([1.5, 1.5], axes=("x", "y"))], + ids=["rotation", "scale"], +) +def test_transform_labels_preserves_label_ids(multiscale: bool, via_spatialdata: bool, transformation): + """Labels are resampled with nearest neighbour, so no label ids are invented (gh-1202).""" + arr = np.zeros((64, 64), dtype=np.uint16) + arr[:32, :] = 1 + arr[32:, :] = 50 + labels = Labels2DModel.parse(arr, scale_factors=[2] if multiscale else None) + set_transformation(labels, transformation, "transformed") + + if via_spatialdata: + transformed = transform(SpatialData(labels={"labels": labels}), to_coordinate_system="transformed")["labels"] + else: + transformed = transform(labels, to_coordinate_system="transformed") + + for ids in _label_ids_per_scale(transformed): + assert ids <= {0, 1, 50} + assert {1, 50} <= ids + + +@pytest.mark.parametrize("multiscale", [False, True]) +def test_transform_labels_3d_preserves_label_ids(multiscale: bool): + """Labels are resampled with nearest neighbour, so no label ids are invented (gh-1202).""" + arr = np.zeros((8, 32, 32), dtype=np.uint16) + arr[:, :16, :] = 1 + arr[:, 16:, :] = 50 + labels = Labels3DModel.parse(arr, scale_factors=[2] if multiscale else None) + set_transformation(labels, Scale([1.5, 1.5, 1.5], axes=("x", "y", "z")), "transformed") + + transformed = transform(labels, to_coordinate_system="transformed") + + for ids in _label_ids_per_scale(transformed): + assert ids <= {0, 1, 50} + assert {1, 50} <= ids + + def test_transform_datatree_scale_handling(): """ Test the cases in which the lowest and highest scale of the result of a