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
1 change: 1 addition & 0 deletions docs/changelog.md
Original file line number Diff line number Diff line change
Expand Up @@ -13,4 +13,5 @@ For developers, please consult the [contributing guide](https://github.com/scver

### Fixed

- `aggregate(values=<image>, by=<multiscale labels>)` 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.
11 changes: 10 additions & 1 deletion src/spatialdata/_core/operations/aggregate.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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:
Expand Down
25 changes: 25 additions & 0 deletions tests/core/operations/test_aggregations.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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:
Expand Down
Loading