diff --git a/CHANGELOG.md b/CHANGELOG.md index f85ea2f7..bbd77451 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -22,6 +22,7 @@ and this project adheres to [Semantic Versioning][]. - Multipolygons are now handled correctly (#93) - Legend order is now deterministic (#143) - Images no longer normalised by default (#150) +- Filtering of shapes and points using the `groups` argument is now possible, coloring by palette and cmap arguments works for shapes and points (#153) - Colorbar no longer autoscales to [0, 1] (#155) ## [0.0.4] - 2023-08-11 diff --git a/src/spatialdata_plot/pl/basic.py b/src/spatialdata_plot/pl/basic.py index e228f2ac..f58d7e04 100644 --- a/src/spatialdata_plot/pl/basic.py +++ b/src/spatialdata_plot/pl/basic.py @@ -14,7 +14,7 @@ from dask.dataframe.core import DataFrame as DaskDataFrame from geopandas import GeoDataFrame from matplotlib.axes import Axes -from matplotlib.colors import Colormap, ListedColormap, Normalize +from matplotlib.colors import Colormap, Normalize from matplotlib.figure import Figure from multiscale_spatial_image.multiscale_spatial_image import MultiscaleSpatialImage from pandas.api.types import is_categorical_dtype @@ -150,7 +150,7 @@ def render_shapes( outline_width: float = 1.5, outline_color: str | list[float] = "#000000ff", layer: str | None = None, - palette: ListedColormap | str | None = None, + palette: str | list[str] | None = None, cmap: Colormap | str | None = None, norm: bool | Normalize = False, na_color: str | tuple[float, ...] | None = "lightgrey", @@ -182,9 +182,13 @@ def render_shapes( layer Key in :attr:`anndata.AnnData.layers` or `None` for :attr:`anndata.AnnData.X`. palette - Palette for discrete annotations, see :class:`matplotlib.colors.Colormap`. + Palette for discrete annotations. List of valid color names that should be used + for the categories (all or as specified by `groups`). For a single category, + a valid color name can be given as string. cmap Colormap for continuous annotations, see :class:`matplotlib.colors.Colormap`. + If no palette is given and `color` refers to a categorical, the colors are + sampled from this colormap. norm Colormap normalization for continuous annotations, see :class:`matplotlib.colors.Normalize`. na_color @@ -235,7 +239,7 @@ def render_points( color: str | None = None, groups: str | Sequence[str] | None = None, size: float = 1.0, - palette: ListedColormap | str | None = None, + palette: str | list[str] | None = None, cmap: Colormap | str | None = None, norm: None | Normalize = None, na_color: str | tuple[float, ...] | None = (0.0, 0.0, 0.0, 0.0), @@ -258,9 +262,13 @@ def render_points( size Value to scale points. palette - Palette for discrete annotations, see :class:`matplotlib.colors.Colormap`. + Palette for discrete annotations. List of valid color names that should be used + for the categories (all or as specified by `groups`). For a single category, + a valid color name can be given as string. cmap Colormap for continuous annotations, see :class:`matplotlib.colors.Colormap`. + If no palette is given and `color` refers to a categorical, the colors are + sampled from this colormap. norm Colormap normalization for continuous annotations, see :class:`matplotlib.colors.Normalize`. na_color @@ -303,7 +311,7 @@ def render_images( cmap: list[Colormap] | list[str] | Colormap | str | None = None, norm: None | Normalize = None, na_color: str | tuple[float, ...] | None = (0.0, 0.0, 0.0, 0.0), - palette: ListedColormap | str | None = None, + palette: str | list[str] | None = None, alpha: float = 1.0, quantiles_for_norm: tuple[float | None, float | None] = (None, None), **kwargs: Any, @@ -381,7 +389,7 @@ def render_labels( contour_px: int = 3, outline: bool = False, layer: str | None = None, - palette: ListedColormap | str | None = None, + palette: str | list[str] | None = None, cmap: Colormap | str | None = None, norm: None | Normalize = None, na_color: str | tuple[float, ...] | None = (0.0, 0.0, 0.0, 0.0), diff --git a/src/spatialdata_plot/pl/render.py b/src/spatialdata_plot/pl/render.py index c849be41..0170f613 100644 --- a/src/spatialdata_plot/pl/render.py +++ b/src/spatialdata_plot/pl/render.py @@ -4,6 +4,7 @@ from copy import copy from typing import Union +import dask import geopandas as gpd import matplotlib import numpy as np @@ -18,6 +19,7 @@ from spatialdata.models import ( Image2DModel, Labels2DModel, + PointsModel, ) from spatialdata_plot._logging import logger @@ -57,6 +59,12 @@ def _render_shapes( ) -> None: elements = render_params.elements + if render_params.groups is not None: + if isinstance(render_params.groups, str): + render_params.groups = [render_params.groups] + if not all(isinstance(g, str) for g in render_params.groups): + raise TypeError("All groups must be strings.") + sdata_filt = sdata.filter_by_coordinate_system( coordinate_system=coordinate_system, filter_table=sdata.table is not None, @@ -68,7 +76,6 @@ def _render_shapes( elements = list(sdata_filt.shapes.keys()) for e in elements: - # shapes = [sdata.shapes[e] for e in elements] shapes = sdata.shapes[e] n_shapes = sum([len(s) for s in shapes]) @@ -88,6 +95,7 @@ def _render_shapes( palette=render_params.palette, na_color=render_params.cmap_params.na_color, alpha=render_params.fill_alpha, + cmap_params=render_params.cmap_params, ) values_are_categorical = color_source_vector is not None @@ -101,7 +109,15 @@ def _render_shapes( if len(color_vector) == 0: color_vector = [render_params.cmap_params.na_color] + # filter by `groups` + if render_params.groups is not None and color_source_vector is not None: + mask = color_source_vector.isin(render_params.groups) + shapes = shapes[mask] + shapes = shapes.reset_index() + color_source_vector = color_source_vector[mask] + color_vector = color_vector[mask] shapes = gpd.GeoDataFrame(shapes, geometry="geometry") + _cax = _get_collection_shape( shapes=shapes, s=render_params.scale, @@ -122,9 +138,12 @@ def _render_shapes( cax = ax.add_collection(_cax) # Using dict.fromkeys here since set returns in arbitrary order - palette = ( - ListedColormap(dict.fromkeys(color_vector)) if render_params.palette is None else render_params.palette - ) + # remove the color of NaN values, else it might be assigned to a category + # order of color in the palette should agree to order of occurence + if color_source_vector is None: + palette = ListedColormap(dict.fromkeys(color_vector)) + else: + palette = ListedColormap(dict.fromkeys(color_vector[~pd.Categorical(color_source_vector).isnull()])) if not ( len(set(color_vector)) == 1 and list(set(color_vector))[0] == to_hex(render_params.cmap_params.na_color) @@ -159,6 +178,12 @@ def _render_points( scalebar_params: ScalebarParams, legend_params: LegendParams, ) -> None: + if render_params.groups is not None: + if isinstance(render_params.groups, str): + render_params.groups = [render_params.groups] + if not all(isinstance(g, str) for g in render_params.groups): + raise TypeError("All groups must be strings.") + elements = render_params.elements sdata_filt = sdata.filter_by_coordinate_system( @@ -178,6 +203,14 @@ def _render_points( color = [render_params.color] if isinstance(render_params.color, str) else render_params.color coords.extend(color) + points = points[coords].compute() + # points[color[0]].cat.set_categories(render_params.groups, inplace=True) + if render_params.groups is not None: + points = points[points[color].isin(render_params.groups).values] + points[color[0]] = points[color[0]].cat.set_categories(render_params.groups) + points = dask.dataframe.from_pandas(points, npartitions=1) + sdata_filt.points[e] = PointsModel.parse(points, coordinates={"x": "x", "y": "y"}) + point_df = points[coords].compute() # we construct an anndata to hack the plotting functions @@ -204,6 +237,7 @@ def _render_points( palette=render_params.palette, na_color=render_params.cmap_params.na_color, alpha=render_params.alpha, + cmap_params=render_params.cmap_params, ) # color_source_vector is None when the values aren't categorical @@ -226,6 +260,11 @@ def _render_points( if not ( len(set(color_vector)) == 1 and list(set(color_vector))[0] == to_hex(render_params.cmap_params.na_color) ): + if color_source_vector is None: + palette = ListedColormap(dict.fromkeys(color_vector)) + else: + palette = ListedColormap(dict.fromkeys(color_vector[~pd.Categorical(color_source_vector).isnull()])) + _ = _decorate_axs( ax=ax, cax=cax, @@ -233,7 +272,7 @@ def _render_points( adata=adata, value_to_plot=render_params.color, color_source_vector=color_source_vector, - palette=render_params.palette, + palette=palette, alpha=render_params.alpha, na_color=render_params.cmap_params.na_color, legend_fontsize=legend_params.legend_fontsize, @@ -415,6 +454,12 @@ def _render_labels( ) -> None: elements = render_params.elements + if render_params.groups is not None: + if isinstance(render_params.groups, str): + render_params.groups = [render_params.groups] + if not all(isinstance(g, str) for g in render_params.groups): + raise TypeError("All groups must be strings.") + sdata_filt = sdata.filter_by_coordinate_system( coordinate_system=coordinate_system, filter_table=sdata.table is not None, @@ -441,7 +486,7 @@ def _render_labels( table = sdata.table[sdata.table.obs[region_key].isin([label_key])] - # get isntance id based on subsetted table + # get instance id based on subsetted table instance_id = table.obs[instance_key].values # get color vector (categorical or continuous) @@ -455,6 +500,7 @@ def _render_labels( palette=render_params.palette, na_color=render_params.cmap_params.na_color, alpha=render_params.fill_alpha, + cmap_params=render_params.cmap_params, ) if (render_params.fill_alpha != render_params.outline_alpha) and render_params.contour_px is not None: diff --git a/src/spatialdata_plot/pl/render_params.py b/src/spatialdata_plot/pl/render_params.py index 9294dc2d..cca7bd58 100644 --- a/src/spatialdata_plot/pl/render_params.py +++ b/src/spatialdata_plot/pl/render_params.py @@ -19,6 +19,7 @@ class CmapParams: cmap: Colormap norm: Normalize na_color: str | tuple[float, ...] = (0.0, 0.0, 0.0, 0.0) + is_default: bool = True @dataclass diff --git a/src/spatialdata_plot/pl/utils.py b/src/spatialdata_plot/pl/utils.py index 981701dd..1f6aff5a 100644 --- a/src/spatialdata_plot/pl/utils.py +++ b/src/spatialdata_plot/pl/utils.py @@ -571,6 +571,7 @@ def _prepare_cmap_norm( vcenter: float | None = None, **kwargs: Any, ) -> CmapParams: + is_default = cmap is None cmap = copy(matplotlib.colormaps[rcParams["image.cmap"] if cmap is None else cmap]) cmap.set_bad("lightgray" if na_color is None else na_color) @@ -583,7 +584,7 @@ def _prepare_cmap_norm( else: norm = TwoSlopeNorm(vmin=vmin, vmax=vmax, vcenter=vcenter) - return CmapParams(cmap, norm, na_color) + return CmapParams(cmap, norm, na_color, is_default) def _set_outline( @@ -745,8 +746,9 @@ def _normalize( def _get_colors_for_categorical_obs( categories: Sequence[str | int], - palette: ListedColormap | str | None = None, + palette: ListedColormap | str | list[str] | None = None, alpha: float = 1.0, + cmap_params: CmapParams | None = None, ) -> list[str]: """ Return a list of colors for a categorical observation. @@ -768,7 +770,9 @@ def _get_colors_for_categorical_obs( # check if default matplotlib palette has enough colors if palette is None: - if len(rcParams["axes.prop_cycle"].by_key()["color"]) >= len_cat: + if cmap_params is not None and not cmap_params.is_default: + palette = cmap_params.cmap + elif len(rcParams["axes.prop_cycle"].by_key()["color"]) >= len_cat: cc = rcParams["axes.prop_cycle"]() palette = [next(cc)["color"] for _ in range(len_cat)] else: @@ -784,12 +788,11 @@ def _get_colors_for_categorical_obs( "input has more than 103 categories. Uniform " "'grey' color will be used for all categories." ) - # otherwise, single chanels turn out grey + # otherwise, single channels turn out grey color_idx = np.linspace(0, 1, len_cat) if len_cat > 1 else [0.7] if isinstance(palette, str): - cmap = plt.get_cmap(palette) - palette = [to_hex(x) for x in cmap(color_idx, alpha=alpha)] + palette = [to_hex(palette)] elif isinstance(palette, list): palette = [to_hex(x) for x in palette] elif isinstance(palette, ListedColormap): @@ -797,7 +800,7 @@ def _get_colors_for_categorical_obs( elif isinstance(palette, LinearSegmentedColormap): palette = [to_hex(palette(x, alpha=alpha)) for x in color_idx] # type: ignore[attr-defined] else: - raise TypeError(f"Palette is {type(palette)} but should be string or `ListedColormap`.") + raise TypeError(f"Palette is {type(palette)} but should be string or list.") return palette[:len_cat] # type: ignore[return-value] @@ -809,9 +812,10 @@ def _set_color_source_vec( element_name: list[str] | str | None = None, layer: str | None = None, groups: Sequence[str] | str | None = None, - palette: ListedColormap | str | None = None, + palette: str | list[str] | None = None, na_color: str | tuple[float, ...] | None = None, alpha: float = 1.0, + cmap_params: CmapParams | None = None, ) -> tuple[ArrayLike | pd.Series | None, ArrayLike, bool]: if value_to_plot is None: color = np.full(len(element), to_hex(na_color)) # type: ignore[arg-type] @@ -836,6 +840,11 @@ def _set_color_source_vec( # numerical case, return early if not is_categorical_dtype(color_source_vector): + if palette is not None: + logging.warning( + "Ignoring categorical palette which is given for a continuous variable. " + "Consider using `cmap` to pass a ColorMap." + ) return None, color_source_vector, False color_source_vector = pd.Categorical(color_source_vector) # convert, e.g., `pd.Series` @@ -843,8 +852,9 @@ def _set_color_source_vec( if groups is not None: color_source_vector = color_source_vector.remove_categories(categories.difference(groups)) + categories = groups - color_map = dict(zip(categories, _get_colors_for_categorical_obs(categories))) + color_map = dict(zip(categories, _get_colors_for_categorical_obs(categories, palette, cmap_params=cmap_params))) # color_map = _get_palette( # adata=adata, cluster_key=value_to_plot, categories=categories, palette=palette, alpha=alpha # ) @@ -918,7 +928,7 @@ def _get_palette( categories: Sequence[Any], adata: AnnData | None = None, cluster_key: None | str = None, - palette: ListedColormap | str | None = None, + palette: ListedColormap | str | list[str] | None = None, alpha: float = 1.0, ) -> Mapping[str, str] | None: if adata is not None and palette is None: @@ -949,11 +959,13 @@ def _get_palette( return {cat: to_hex(to_rgba(col)[:3]) for cat, col in zip(categories, palette[:len_cat])} if isinstance(palette, str): - cmap = plt.get_cmap(palette) + cmap = ListedColormap([palette]) + elif isinstance(palette, list): + cmap = ListedColormap(palette) elif isinstance(palette, ListedColormap): cmap = palette else: - raise TypeError(f"Palette is {type(palette)} but should be string or `ListedColormap`.") + raise TypeError(f"Palette is {type(palette)} but should be string or list.") palette = [to_hex(np.round(x, 5)) for x in cmap(np.linspace(0, 1, len_cat), alpha=alpha)] return dict(zip(categories, palette)) @@ -970,6 +982,8 @@ def _maybe_set_colors( raise KeyError("Unable to copy the palette when there was other explicitly specified.") target.uns[color_key] = source.uns[color_key] except KeyError: + if isinstance(palette, str): + palette = ListedColormap([palette]) if isinstance(palette, ListedColormap): # `scanpy` requires it palette = cycler(color=palette.colors) add_colors_for_categorical_sample_annotation(target, key=key, force_update_colors=True, palette=palette) @@ -982,7 +996,7 @@ def _decorate_axs( adata: AnnData, value_to_plot: str | None, color_source_vector: pd.Series[CategoricalDtype], - palette: ListedColormap | str | None = None, + palette: ListedColormap | str | list[str] | None = None, alpha: float = 1.0, na_color: str | tuple[float, ...] = (0.0, 0.0, 0.0, 0.0), legend_fontsize: int | float | _FontSize | None = None, @@ -1006,7 +1020,9 @@ def _decorate_axs( # Adding legends if is_categorical_dtype(color_source_vector): - clusters = color_source_vector.categories + # order of clusters should agree to palette order + clusters = color_source_vector.unique() + clusters = clusters[~clusters.isnull()] palette = _get_palette( adata=adata, cluster_key=value_to_plot, categories=clusters, palette=palette, alpha=alpha ) diff --git a/tests/_images/Points_can_filter_with_groups.png b/tests/_images/Points_can_filter_with_groups.png new file mode 100644 index 00000000..3d52aebb Binary files /dev/null and b/tests/_images/Points_can_filter_with_groups.png differ diff --git a/tests/_images/Points_coloring_with_cmap.png b/tests/_images/Points_coloring_with_cmap.png new file mode 100644 index 00000000..4f620cca Binary files /dev/null and b/tests/_images/Points_coloring_with_cmap.png differ diff --git a/tests/_images/Points_coloring_with_palette.png b/tests/_images/Points_coloring_with_palette.png new file mode 100644 index 00000000..6dfaab73 Binary files /dev/null and b/tests/_images/Points_coloring_with_palette.png differ diff --git a/tests/_images/Shapes_can_filter_with_groups.png b/tests/_images/Shapes_can_filter_with_groups.png new file mode 100644 index 00000000..a7cf53c1 Binary files /dev/null and b/tests/_images/Shapes_can_filter_with_groups.png differ diff --git a/tests/_images/Shapes_coloring_with_palette.png b/tests/_images/Shapes_coloring_with_palette.png new file mode 100644 index 00000000..dacb1d6c Binary files /dev/null and b/tests/_images/Shapes_coloring_with_palette.png differ diff --git a/tests/pl/test_render_points.py b/tests/pl/test_render_points.py index bfa58658..43379517 100644 --- a/tests/pl/test_render_points.py +++ b/tests/pl/test_render_points.py @@ -21,3 +21,12 @@ class TestPoints(PlotTester, metaclass=PlotTesterMeta): def test_plot_can_render_points(self, sdata_blobs: SpatialData): sdata_blobs.pl.render_points(elements="blobs_points").pl.show() + + def test_plot_can_filter_with_groups(self, sdata_blobs: SpatialData): + sdata_blobs.pl.render_points(color="genes", groups="b", palette="orange").pl.show() + + def test_plot_coloring_with_palette(self, sdata_blobs: SpatialData): + sdata_blobs.pl.render_points(color="genes", groups=["a", "b"], palette=["lightgreen", "darkblue"]).pl.show() + + def test_plot_coloring_with_cmap(self, sdata_blobs: SpatialData): + sdata_blobs.pl.render_points(color="genes", cmap="rainbow").pl.show() diff --git a/tests/pl/test_render_shapes.py b/tests/pl/test_render_shapes.py index 0a703b2f..3550a6a1 100644 --- a/tests/pl/test_render_shapes.py +++ b/tests/pl/test_render_shapes.py @@ -101,6 +101,26 @@ def test_plot_can_color_from_geodataframe(self, sdata_blobs: SpatialData): def test_plot_can_scale_shapes(self, sdata_blobs: SpatialData): sdata_blobs.pl.render_shapes(elements="blobs_circles", scale=0.5).pl.show() + def test_plot_can_filter_with_groups(self, sdata_blobs: SpatialData): + sdata_blobs.shapes["blobs_polygons"]["cluster"] = "c1" + sdata_blobs.shapes["blobs_polygons"].iloc[3:5, 1] = "c2" + sdata_blobs.shapes["blobs_polygons"]["cluster"] = sdata_blobs.shapes["blobs_polygons"]["cluster"].astype( + "category" + ) + + sdata_blobs.pl.render_shapes("blobs_polygons", color="cluster", groups="c1").pl.show() + + def test_plot_coloring_with_palette(self, sdata_blobs: SpatialData): + sdata_blobs.shapes["blobs_polygons"]["cluster"] = "c1" + sdata_blobs.shapes["blobs_polygons"].iloc[3:5, 1] = "c2" + sdata_blobs.shapes["blobs_polygons"]["cluster"] = sdata_blobs.shapes["blobs_polygons"]["cluster"].astype( + "category" + ) + + sdata_blobs.pl.render_shapes( + "blobs_polygons", color="cluster", groups=["c2", "c1"], palette=["green", "yellow"] + ).pl.show() + def test_plot_colorbar_respects_input_limits(self, sdata_blobs: SpatialData): sdata_blobs.shapes["blobs_polygons"]["cluster"] = [1, 2, 3, 5, 20] sdata_blobs.pl.render_shapes("blobs_polygons", color="cluster", groups=["c1"]).pl.show()