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
48 changes: 44 additions & 4 deletions src/spatialdata/_core/query/relational_query.py
Original file line number Diff line number Diff line change
Expand Up @@ -236,6 +236,9 @@ def _get_masked_element(
mask_values = left_index[mask]
else:
mask_values = left_index
elif mask_values is not None:
order_mask = np.isin(element_indices, mask_values)
mask_values = np.asarray(element_indices)[order_mask]

if isinstance(element, DaskDataFrame):
return element.map_partitions(lambda df: df.loc[mask_values], meta=element)
Expand All @@ -252,6 +255,13 @@ def _right_exclusive_join_spatialelement_table(
match_rows: Literal["left", "no", "right"],
filter_label_pixels: bool | None = None,
) -> tuple[dict[str, Any], AnnData | None]:
if match_rows == "left":
warnings.warn(
"Matching rows 'left' is not supported for 'right_exclusive' join; it will be treated as 'no'.",
UserWarning,
stacklevel=2,
)
match_rows = "no"
regions, region_column_name, instance_key = get_table_keys(table)
if isinstance(regions, str):
regions = [regions]
Expand Down Expand Up @@ -297,7 +307,12 @@ def _right_join_spatialelement_table(
filter_label_pixels: bool | None = None,
) -> tuple[dict[str, Any], AnnData]:
if match_rows == "left":
warnings.warn("Matching rows 'left' is not supported for 'right' join.", UserWarning, stacklevel=2)
warnings.warn(
"Matching rows 'left' is not supported for 'right' join; it will be treated as 'no'.",
UserWarning,
stacklevel=2,
)
match_rows = "no"
regions, region_column_name, instance_key = get_table_keys(table)
if isinstance(regions, str):
regions = [regions]
Expand Down Expand Up @@ -383,6 +398,14 @@ def _inner_join_spatialelement_table(

if joined_indices is not None:
joined_indices = joined_indices.dropna() if any(joined_indices.isna()) else joined_indices
# `groupby(region)` above collects the matching table rows grouped by region, which does not
# preserve the original `table.obs` row order when a table annotates multiple interleaved
# regions. For `match_rows="no"` there is no element-driven ordering to honor, and for
# `match_rows="right"` the table's own row order takes priority (only `match_rows="left"` lets the
# element's row order override it), so in both cases restore the original table row order, as
# would be expected for a semi-join.
if match_rows in ("no", "right"):
joined_indices = joined_indices.sort_values()

joined_table = table[joined_indices.tolist(), :].copy() if joined_indices is not None else None

Expand All @@ -401,6 +424,13 @@ def _left_exclusive_join_spatialelement_table(
match_rows: Literal["left", "no", "right"],
filter_label_pixels: bool | None = None,
) -> tuple[dict[str, Any], AnnData | None]:
if match_rows == "right":
warnings.warn(
"Matching rows 'right' is not supported for 'left_exclusive' join; it will be treated as 'no'.",
UserWarning,
stacklevel=2,
)
match_rows = "no"
regions, region_column_name, instance_key = get_table_keys(table)
if isinstance(regions, str):
regions = [regions]
Expand All @@ -411,8 +441,7 @@ def _left_exclusive_join_spatialelement_table(
group_df = groups_df.get_group(name)
table_instance_key_column = group_df[instance_key]
if element_type in ["points", "shapes"]:
mask = np.full(len(element), True, dtype=bool)
mask[table_instance_key_column.values] = False
mask = ~np.isin(element.index, table_instance_key_column.values)
masked_element = element.loc[mask, :] if mask.sum() != 0 else None
element_dict[element_type][name] = masked_element
else:
Expand All @@ -438,7 +467,12 @@ def _left_join_spatialelement_table(
filter_label_pixels: bool | None = None,
) -> tuple[dict[str, Any], AnnData]:
if match_rows == "right":
warnings.warn("Matching rows 'right' is not supported for 'left' join.", UserWarning, stacklevel=2)
warnings.warn(
"Matching rows 'right' is not supported for 'left' join; it will be treated as 'no'.",

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Truth be told, for left join we could support the right match_rows, and for right join we could support the left match_rows. But we can skip it for now (since also it was not supported before this PR), and eventually do it in the future.

For left_exclusive, right match_rows does not make sense, so it is not supported. Same for right_exclusive: left match_rows does not make sense there.

UserWarning,
stacklevel=2,
)
match_rows = "no"
regions, region_column_name, instance_key = get_table_keys(table)
if isinstance(regions, str):
regions = [regions]
Expand Down Expand Up @@ -469,6 +503,12 @@ def _left_join_spatialelement_table(
# if nan were present, the dtype would have been changed to float
if joined_indices.dtype == float:
joined_indices = joined_indices.astype(int)
# `groupby(region)` above collects the matching table rows grouped by region, which does not
# preserve the original `table.obs` row order when a table annotates multiple interleaved
# regions. For `match_rows="no"` there is no element-driven ordering to honor, so
# restore the original table row order, as would be expected for a semi-join.
if match_rows == "no":
joined_indices = joined_indices.sort_values()
joined_table = table[joined_indices.tolist(), :].copy() if joined_indices is not None else None
_inplace_fix_subset_categorical_obs(subset_adata=joined_table, original_adata=table)
if joined_table is not None:
Expand Down
164 changes: 164 additions & 0 deletions tests/core/query/test_relational_query.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
from __future__ import annotations

import re
import warnings
from dataclasses import dataclass, field

import annsel as an
import numpy as np
import pandas as pd
Expand Down Expand Up @@ -940,6 +944,166 @@ def test_filter_table_categorical_bug(shapes):
shapes.filter_by_coordinate_system("global")


@dataclass
class _JoinOutcome:
"""Expected outcome of a `join_spatialelement_table()` call, for a given `how`/`match_rows` pair."""

# expected `joined_table.obs["label"]`, in order; `None` when no table is expected to be returned
table_order: list[str] | None = None
# whether a "Matching rows '<...>' is not supported for '<...>' join." UserWarning is expected to be emitted
warns: bool = False
# expected values of `element_dict[name].index` for element name in {"a", "b"}; `None` for a element name whose
# returned join result is expected to be `None` (e.g. fully excluded, or not returned by this join type)
element_index: dict[str, list[int] | None] = field(default_factory=dict)


def _make_interleaved_regions_sdata() -> tuple[SpatialData, dict[str, dict[str, _JoinOutcome]]]:
from geopandas import GeoDataFrame
from shapely.geometry import Point

from spatialdata.models import ShapesModel

def circles(indices):
# `indices` gives both the number of circles and the (non-default) row order of the element.
gdf = GeoDataFrame(
{"geometry": [Point(i, i) for i in range(len(indices))], "radius": [1.0] * len(indices)},
index=pd.Index(indices),
)
return ShapesModel.parse(gdf)

# assumptions/comments:
# - no duplicate values in the index of each spatial element (duplicate values are tested elsewhere)
# - no duplicate values for the instance_key column of the table (duplicate values are tested elsewhere)
#
# edge cases being tested:
# - instance_id values are non-monotonic
# - the index in each spatial element is non-monotonic
# - we also set the index of the table obs to random values; these should be ignored (in the code we call .index on
# a region_key column, but the index is freshly reset by a nearby call of .reset_index() inside the join
# machinery)
obs = pd.DataFrame(
{
"region": pd.Categorical(["b", "b", "a", "b", "a", "a", "b"]),
"instance_id": [2, 1, 2, 3, 1, 0, 0],
"label": ["b2", "b1", "a2", "b3", "a1", "a0", "b0"],
},
index=np.random.default_rng(0).integers(0, 3, size=7).astype(str),
)
shapes = {"a": circles([2, 1, 0]), "b": circles([1, 2, 0])}
# to make understanding easier, you may want to refer to the figure on joins from the docs:
# https://spatialdata.scverse.org/en/stable/tutorials/notebooks/notebooks/examples/tables.html
expected = {
"left": {
"no": _JoinOutcome(
table_order=["b2", "b1", "a2", "a1", "a0", "b0"], element_index={"a": [2, 1, 0], "b": [1, 2, 0]}
),
"left": _JoinOutcome(
table_order=["a2", "a1", "a0", "b1", "b2", "b0"], element_index={"a": [2, 1, 0], "b": [1, 2, 0]}
),
"right": _JoinOutcome(
table_order=["b2", "b1", "a2", "a1", "a0", "b0"],
warns=True,
element_index={"a": [2, 1, 0], "b": [1, 2, 0]},
),
},
"left_exclusive": {
# TODO: make this test more interesting by adding indices 5, 4 to "a" and 4, 6 to "b"
# by design, "left_exclusive" never returns a table (only filtered elements), regardless of
# match_rows or whether anything was actually excluded.
Comment on lines +1010 to +1012

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We can do in this PR or leave for the future. It is beyond the scope of the original fix anyway.

"no": _JoinOutcome(table_order=None, element_index={"a": None, "b": None}),
"left": _JoinOutcome(table_order=None, element_index={"a": None, "b": None}),
"right": _JoinOutcome(table_order=None, warns=True, element_index={"a": None, "b": None}),
},
"inner": {
"no": _JoinOutcome(
table_order=["b2", "b1", "a2", "a1", "a0", "b0"],
element_index={"a": [2, 1, 0], "b": [1, 2, 0]},
),
"left": _JoinOutcome(
table_order=["a2", "a1", "a0", "b1", "b2", "b0"], element_index={"a": [2, 1, 0], "b": [1, 2, 0]}
),
"right": _JoinOutcome(
table_order=["b2", "b1", "a2", "a1", "a0", "b0"], element_index={"a": [2, 1, 0], "b": [2, 1, 0]}
),
},
"right": {
"no": _JoinOutcome(
table_order=["b2", "b1", "a2", "b3", "a1", "a0", "b0"],
element_index={"a": [2, 1, 0], "b": [1, 2, 0]},
),
"left": _JoinOutcome(
table_order=["b2", "b1", "a2", "b3", "a1", "a0", "b0"],
warns=True,
element_index={"a": [2, 1, 0], "b": [1, 2, 0]},
),
"right": _JoinOutcome(
table_order=["b2", "b1", "a2", "b3", "a1", "a0", "b0"], element_index={"a": [2, 1, 0], "b": [2, 1, 0]}
),
},
"right_exclusive": {
"no": _JoinOutcome(table_order=["b3"], element_index={"a": None, "b": None}),
"left": _JoinOutcome(table_order=["b3"], warns=True, element_index={"a": None, "b": None}),
"right": _JoinOutcome(table_order=["b3"], element_index={"a": None, "b": None}),
},
}

table = TableModel.parse(
AnnData(X=np.zeros((len(obs), 1)), obs=obs),
region=["a", "b"],
region_key="region",
instance_key="instance_id",
)
sdata = SpatialData(shapes=shapes, tables={"table": table})
return sdata, expected


@pytest.mark.parametrize("match_rows", ["no", "left", "right"])
@pytest.mark.parametrize("how", ["left", "left_exclusive", "inner", "right", "right_exclusive"])
def test_join_preserves_row_order_multiple_interleaved_regions(how, match_rows):
# generalization to all the join types of the bug reported in https://github.com/scverse/spatialdata/issues/1162
# covering all `how` values of `join_spatialelement_table`, crossed with all values of `match_rows`, and checking
# whether the row orders of the returned spatial elements and table are correct and if the "match_rows not
# supported" UserWarning is (or isn't) actually raised (see `_make_interleaved_regions_sdata` and `_JoinOutcome`).
sdata, expected_by_how_and_match_rows = _make_interleaved_regions_sdata()
outcome = expected_by_how_and_match_rows[how][match_rows]

with warnings.catch_warnings(record=True) as record:
warnings.simplefilter("always")
element_dict, joined_table = join_spatialelement_table(
sdata=sdata,
spatial_element_names=["a", "b"],
table_name="table",
how=how,
match_rows=match_rows,
)

# other UserWarnings can also fire here (e.g. anndata's "Observation names are not unique", triggered by
# the fixture's duplicated obs_names), so only look for the one this test is actually about. The message
# looks like "Matching rows 'right' is not supported for 'left_exclusive' join; it will be treated as 'no'.",
# with the two quoted values varying by `match_rows` / `how`.
unsupported_match_rows_re = re.compile(
r"Matching rows '[^']+' is not supported for '[^']+' join; it will be treated as 'no'\."
)
unsupported_match_rows_warnings = [
w for w in record if issubclass(w.category, UserWarning) and unsupported_match_rows_re.search(str(w.message))
]
assert bool(unsupported_match_rows_warnings) == outcome.warns

if outcome.table_order is None:
assert joined_table is None
else:
assert joined_table is not None
assert list(joined_table.obs["label"]) == outcome.table_order

for name, expected_index in outcome.element_index.items():
actual_element = element_dict[name]
if expected_index is None:
assert actual_element is None
else:
assert actual_element is not None
assert list(actual_element.index) == expected_index


def test_filter_table_non_annotating(full_sdata):

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Test LGTM

Copy link
Copy Markdown
Member Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I expanded them, please can you re-review?

obs = pd.DataFrame({"test": ["a", "b", "c"]}, index=list(map(str, range(3))))
adata = AnnData(obs=obs)
Expand Down