diff --git a/src/squidpy/gr/neighbors.py b/src/squidpy/gr/neighbors.py index 5a2bb7514..865be7372 100644 --- a/src/squidpy/gr/neighbors.py +++ b/src/squidpy/gr/neighbors.py @@ -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, @@ -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 { @@ -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 { diff --git a/tests/graph/test_spatial_neighbors.py b/tests/graph/test_spatial_neighbors.py index 63b441c84..a82fca037 100644 --- a/tests/graph/test_spatial_neighbors.py +++ b/tests/graph/test_spatial_neighbors.py @@ -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, @@ -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)