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
11 changes: 8 additions & 3 deletions src/squidpy/gr/neighbors.py
Original file line number Diff line number Diff line change
Expand Up @@ -88,7 +88,10 @@ def postprocessors(self) -> Sequence[GraphPostprocessor[GraphMatrixT]]:

@abstractmethod
def uns_params(self) -> dict[str, Any]:
"""Parameters stored in :attr:`anndata.AnnData.uns` after graph construction."""
"""Parameters stored in :attr:`anndata.AnnData.uns` after graph construction.

Values must be writable by :mod:`anndata`, e.g. a :class:`list`, not a :class:`tuple`.
"""

def combine(
self,
Expand Down Expand Up @@ -234,7 +237,8 @@ def __init__(
percentile=percentile,
postprocessors=postprocessors,
)
self.radius = radius
# Store intervals as a list: :mod:`anndata` cannot write tuples to ``uns``.
self.radius = list(radius) if isinstance(radius, tuple) else radius

def uns_params(self) -> dict[str, Any]:
return {
Expand Down Expand Up @@ -303,7 +307,8 @@ def __init__(
percentile=percentile,
postprocessors=postprocessors,
)
self.radius = radius
# Store intervals as a list: :mod:`anndata` cannot write tuples to ``uns``.
self.radius = list(radius) if isinstance(radius, tuple) else radius

def uns_params(self) -> dict[str, Any]:
return {
Expand Down
25 changes: 24 additions & 1 deletion tests/graph/test_spatial_neighbors.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,7 +12,13 @@

from squidpy._constants._constants import Transform
from squidpy._constants._pkg_constants import Key
from squidpy.gr import mask_graph, spatial_neighbors, spatial_neighbors_from_builder
from squidpy.gr import (
mask_graph,
spatial_neighbors,
spatial_neighbors_delaunay,
spatial_neighbors_from_builder,
spatial_neighbors_radius,
)
from squidpy.gr.neighbors import (
DelaunayBuilder,
GridBuilder,
Expand Down Expand Up @@ -277,6 +283,23 @@ def test_delaunay_builder_scalar_radius_equals_zero_max_tuple(self, non_visium_a
np.testing.assert_array_equal(scalar.connectivities.toarray(), interval.connectivities.toarray())
np.testing.assert_allclose(scalar.distances.toarray(), interval.distances.toarray())

@pytest.mark.parametrize(
("func", "radius", "expected"),
[
(spatial_neighbors_radius, 5.0, 5.0),
(spatial_neighbors_radius, (2.0, 4.0), [2.0, 4.0]),
(spatial_neighbors_delaunay, (2.0, 4.0), [2.0, 4.0]),
(spatial_neighbors_delaunay, 5.0, [0.0, 5.0]),
],
ids=["radius_scalar", "radius_interval", "delaunay_interval", "delaunay_scalar"],
)
def test_radius_stored_in_uns_is_writable(self, non_visium_adata: AnnData, tmp_path, func, radius, expected):
func(non_visium_adata, radius=radius)

# an interval radius is stored as a list: `anndata` cannot write tuples
assert non_visium_adata.uns[Key.uns.spatial_neighs()]["params"]["radius"] == expected
non_visium_adata.write_h5ad(tmp_path / "adata.h5ad")

def test_delaunay_mode_warns_on_n_neighs(self, non_visium_adata: AnnData):
with pytest.warns(FutureWarning, match=r"Parameter `n_neighs` is ignored when `delaunay=True`"):
spatial_neighbors(non_visium_adata, coord_type="generic", delaunay=True, n_neighs=3, copy=True)
Expand Down
Loading