Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 3 additions & 3 deletions src/spatialdata/_core/operations/transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 = {}
Expand Down
54 changes: 53 additions & 1 deletion tests/core/operations/test_transform.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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
Expand Down
Loading