From ea9069e20b52724db7fba29309722d5b6f133900 Mon Sep 17 00:00:00 2001 From: Sizerta Date: Thu, 24 Sep 2026 15:46:51 +0400 Subject: [PATCH] fix: aggregate images by multiscale labels --- docs/changelog.md | 1 + src/spatialdata/_core/operations/aggregate.py | 11 +++++++- tests/core/operations/test_aggregations.py | 25 +++++++++++++++++++ 3 files changed, 36 insertions(+), 1 deletion(-) diff --git a/docs/changelog.md b/docs/changelog.md index 89658c92a..4b6c2186f 100644 --- a/docs/changelog.md +++ b/docs/changelog.md @@ -13,4 +13,5 @@ For developers, please consult the [contributing guide](https://github.com/scver ### Fixed +- `aggregate(values=, by=)` no longer crashes with `AttributeError: 'DataTree' object has no attribute 'dtype'`. Multiscale labels passed as `by` now aggregate to the same table as their full-resolution counterpart. - Querying a `SpatialData` object with multiple (batched) bounding boxes now raises an explicit `NotImplementedError` instead of failing with an opaque `AssertionError`. Batched queries are supported when querying a `SpatialElement` directly, but not when querying a `SpatialData` object, since that would require returning one `SpatialData` object per bounding box. diff --git a/src/spatialdata/_core/operations/aggregate.py b/src/spatialdata/_core/operations/aggregate.py index b4772df1b..7bc686606 100644 --- a/src/spatialdata/_core/operations/aggregate.py +++ b/src/spatialdata/_core/operations/aggregate.py @@ -20,6 +20,7 @@ from spatialdata._core.query.relational_query import get_values from spatialdata._core.spatialdata import SpatialData from spatialdata._types import ArrayLike +from spatialdata._utils import get_pyramid_levels from spatialdata.models import Image2DModel, Labels2DModel, PointsModel, ShapesModel, TableModel, get_model from spatialdata.transformations import BaseTransformation, Identity, get_transformation @@ -230,7 +231,15 @@ def _create_sdata_from_table_and_regions( ) -> SpatialData: from spatialdata._core._deepcopy import deepcopy as _deepcopy - shapes_index_dtype = shapes.index.dtype if isinstance(shapes, GeoDataFrame) else shapes.dtype + if isinstance(shapes, GeoDataFrame): + shapes_index_dtype = shapes.index.dtype + elif isinstance(shapes, DataTree): + # multiscale labels: every scale has the dtype of the full-resolution one + scale0 = get_pyramid_levels(shapes, n=0) + assert isinstance(scale0, DataArray) + shapes_index_dtype = scale0.dtype + else: + shapes_index_dtype = shapes.dtype try: table.obs[instance_key] = table.obs_names.copy().astype(shapes_index_dtype) except ValueError as err: diff --git a/tests/core/operations/test_aggregations.py b/tests/core/operations/test_aggregations.py index d471b79cc..e8c45cfc5 100644 --- a/tests/core/operations/test_aggregations.py +++ b/tests/core/operations/test_aggregations.py @@ -8,6 +8,7 @@ from anndata.tests.helpers import assert_equal from geopandas import GeoDataFrame from numpy.random import default_rng +from scipy.sparse import issparse from spatialdata import aggregate, to_polygons from spatialdata._core._deepcopy import deepcopy as _deepcopy @@ -359,6 +360,30 @@ def test_aggregate_image_by_labels(labels_blobs, image_schema, labels_schema) -> assert len(out) == 3 +@pytest.mark.parametrize("multiscale_values", [False, True]) +@pytest.mark.parametrize("multiscale_by", [False, True]) +def test_aggregate_image_by_labels_multiscale(labels_blobs, *, multiscale_values, multiscale_by) -> None: + """Multiscale `values` and `by` aggregate like their single-scale counterparts.""" + image = RNG.normal(size=(3,) + labels_blobs.shape) + scale_factors = [2] + + single = aggregate( + values=Image2DModel.parse(image), + by=Labels2DModel.parse(labels_blobs), + agg_func="mean", + ).tables["table"] + out = aggregate( + values=Image2DModel.parse(image, scale_factors=scale_factors if multiscale_values else None), + by=Labels2DModel.parse(labels_blobs, scale_factors=scale_factors if multiscale_by else None), + agg_func="mean", + ).tables["table"] + + assert len(out) + 1 == len(np.unique(labels_blobs)) + np.testing.assert_array_equal(out.obs_names, single.obs_names) + x, x_single = (a.toarray() if issparse(a) else a for a in (out.X, single.X)) + np.testing.assert_allclose(x, x_single) + + @pytest.mark.parametrize("values", ["blobs_image", "blobs_points", "blobs_circles", "blobs_polygons"]) @pytest.mark.parametrize("by", ["blobs_labels", "blobs_circles", "blobs_polygons"]) def test_aggregate_requiring_alignment(sdata_blobs: SpatialData, values, by) -> None: