diff --git a/CHANGELOG.md b/CHANGELOG.md index 9e75c071..a87fd0ab 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -8,15 +8,17 @@ and this project adheres to [Semantic Versioning][]. [keep a changelog]: https://keepachangelog.com/en/1.0.0/ [semantic versioning]: https://semver.org/spec/v2.0.0.html -## [0.0.5] -tbd +## [0.1.0] - tbd ### Added - Multipolygons are now handled correctly (#93) +- Can now plot columns from GeoDataFrame (#149) ### Fixed - Legend order is now deterministic (#143) +- Images no longer normalised by default (#150) ## [0.0.4] - 2023-08-11 diff --git a/src/spatialdata_plot/pl/basic.py b/src/spatialdata_plot/pl/basic.py index 51839213..795e5ec7 100644 --- a/src/spatialdata_plot/pl/basic.py +++ b/src/spatialdata_plot/pl/basic.py @@ -23,20 +23,22 @@ from spatialdata_plot._accessor import register_spatial_data_accessor from spatialdata_plot.pl.render import ( - ImageRenderParams, - LabelsRenderParams, - PointsRenderParams, - ShapesRenderParams, _render_images, _render_labels, _render_points, _render_shapes, ) -from spatialdata_plot.pl.utils import ( +from spatialdata_plot.pl.render_params import ( CmapParams, + ImageRenderParams, + LabelsRenderParams, LegendParams, + PointsRenderParams, + ShapesRenderParams, _FontSize, _FontWeight, +) +from spatialdata_plot.pl.utils import ( _get_cs_contents, _get_extent, _maybe_set_colors, @@ -147,7 +149,6 @@ def render_shapes( outline: bool = False, outline_width: float = 1.5, outline_color: str | list[float] = "#000000ff", - alt_var: str | None = None, layer: str | None = None, palette: ListedColormap | str | None = None, cmap: Colormap | str | None = None, @@ -178,8 +179,6 @@ def render_shapes( Width of the border. outline_color Color of the border. - alt_var - Which column to use in :attr:`anndata.AnnData.var` to select alternative ``var_name``. layer Key in :attr:`anndata.AnnData.layers` or `None` for :attr:`anndata.AnnData.X`. palette @@ -219,7 +218,6 @@ def render_shapes( color=color, groups=groups, outline_params=outline_params, - alt_var=alt_var, layer=layer, cmap_params=cmap_params, palette=palette, @@ -381,7 +379,6 @@ def render_labels( groups: str | Sequence[str] | None = None, contour_px: int = 3, outline: bool = False, - alt_var: str | None = None, layer: str | None = None, palette: ListedColormap | str | None = None, cmap: Colormap | str | None = None, @@ -409,8 +406,6 @@ def render_labels( entire segment, see :func:`skimage.morphology.erosion`. outline Whether to plot boundaries around segmentation masks. - alt_var - Which column to use in :attr:`anndata.AnnData.var` to select alternative ``var_name``. layer Key in :attr:`anndata.AnnData.layers` or `None` for :attr:`anndata.AnnData.X`. palette @@ -452,7 +447,6 @@ def render_labels( groups=groups, contour_px=contour_px, outline=outline, - alt_var=alt_var, layer=layer, cmap_params=cmap_params, palette=palette, @@ -667,15 +661,6 @@ def show( # extent=extent[cs], ) elif cmd == "render_shapes" and cs_contents.query(f"cs == '{cs}'")["has_shapes"][0]: - if sdata.table is not None and isinstance(params.color, str): - colors = sc.get.obs_df(sdata.table, params.color) - if is_categorical_dtype(colors): - _maybe_set_colors( - source=sdata.table, - target=sdata.table, - key=params.color, - palette=params.palette, - ) _render_shapes( sdata=sdata, render_params=params, @@ -728,6 +713,7 @@ def show( else: t = cs ax.set_title(t) + ax.set_aspect("equal") if any( [ diff --git a/src/spatialdata_plot/pl/render.py b/src/spatialdata_plot/pl/render.py index 62ad0dce..f5ceac11 100644 --- a/src/spatialdata_plot/pl/render.py +++ b/src/spatialdata_plot/pl/render.py @@ -2,9 +2,7 @@ from collections.abc import Sequence from copy import copy -from dataclasses import dataclass -from functools import partial -from typing import Any, Callable, Union +from typing import Union import geopandas as gpd import matplotlib @@ -14,11 +12,7 @@ import spatial_image import spatialdata as sd from anndata import AnnData -from geopandas import GeoDataFrame -from matplotlib import colors -from matplotlib.collections import PatchCollection -from matplotlib.colors import ColorConverter, ListedColormap, Normalize -from matplotlib.patches import Circle, Polygon +from matplotlib.colors import ListedColormap, Normalize from pandas.api.types import is_categorical_dtype from scanpy._settings import settings as sc_settings from spatialdata.models import ( @@ -27,44 +21,29 @@ ) from spatialdata_plot._logging import logger -from spatialdata_plot.pl.utils import ( - CmapParams, +from spatialdata_plot.pl.render_params import ( FigParams, + ImageRenderParams, + LabelsRenderParams, LegendParams, - OutlineParams, + PointsRenderParams, ScalebarParams, + ShapesRenderParams, +) +from spatialdata_plot.pl.utils import ( _decorate_axs, + _get_collection_shape, _get_colors_for_categorical_obs, _get_linear_colormap, - _make_patch_from_multipolygon, _map_color_seg, _maybe_set_colors, _normalize, _set_color_source_vec, + to_hex, ) from spatialdata_plot.pp.utils import _get_instance_key, _get_region_key _Normalize = Union[Normalize, Sequence[Normalize]] -to_hex = partial(colors.to_hex, keep_alpha=True) - - -@dataclass -class ShapesRenderParams: - """Labels render parameters..""" - - cmap_params: CmapParams - outline_params: OutlineParams - elements: str | Sequence[str] | None = None - color: str | None = None - groups: str | Sequence[str] | None = None - contour_px: int | None = None - alt_var: str | None = None - layer: str | None = None - palette: ListedColormap | str | None = None - outline_alpha: float = 1.0 - fill_alpha: float = 0.3 - size: float = 1.0 - transfunc: Callable[[float], float] | None = None def _render_shapes( @@ -88,190 +67,83 @@ def _render_shapes( if elements is None: elements = list(sdata_filt.shapes.keys()) - shapes = [sdata.shapes[e] for e in elements] - n_shapes = sum([len(s) for s in shapes]) - - if sdata.table is None: - table = AnnData(None, obs=pd.DataFrame(index=pd.Index(np.arange(n_shapes), dtype=str))) - else: - table = sdata.table[sdata.table.obs[_get_region_key(sdata)].isin(elements)] - - # get color vector (categorical or continuous) - color_source_vector, color_vector, _ = _set_color_source_vec( - adata=table, - value_to_plot=render_params.color, - alt_var=render_params.alt_var, - layer=render_params.layer, - groups=render_params.groups, - palette=render_params.palette, - na_color=render_params.cmap_params.na_color, - alpha=render_params.fill_alpha, - ) - - # color_source_vector is None when the values aren't categorical - if color_source_vector is None and render_params.transfunc is not None: - color_vector = render_params.transfunc(color_vector) - - def _get_collection_shape( - shapes: list[GeoDataFrame], - c: Any, - s: float, - norm: Any, - fill_alpha: None | float = None, - outline_alpha: None | float = None, - **kwargs: Any, - ) -> PatchCollection: - """ - Get a PatchCollection for rendering given geometries with specified colors and outlines. - - Args: - - shapes (list[GeoDataFrame]): List of geometrical shapes. - - c: Color parameter. - - s (float): Size of the shape. - - norm: Normalization for the color map. - - fill_alpha (float, optional): Opacity for the fill color. - - outline_alpha (float, optional): Opacity for the outline. - - **kwargs: Additional keyword arguments. - - Returns - ------- - - PatchCollection: Collection of patches for rendering. - """ - cmap = kwargs["cmap"] - - try: - # fails when numeric - fill_c = ColorConverter().to_rgba_array(c) - except ValueError: - if norm is None: - c = cmap(c) - else: - norm = colors.Normalize(vmin=min(c), vmax=max(c)) - c = cmap(norm(c)) - - fill_c = ColorConverter().to_rgba_array(c) - fill_c[..., -1] = render_params.fill_alpha + 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]) - if render_params.outline_params.outline: - outline_c = ColorConverter().to_rgba_array(render_params.outline_params.outline_color) - outline_c[..., -1] = render_params.outline_alpha - outline_c = outline_c.tolist() + if sdata.table is None: + table = AnnData(None, obs=pd.DataFrame(index=pd.Index(np.arange(n_shapes), dtype=str))) else: - outline_c = [None] - outline_c = outline_c * fill_c.shape[0] - - shapes_df = pd.DataFrame(shapes, copy=True) - - # remove empty points/polygons - shapes_df = shapes_df[shapes_df["geometry"].apply(lambda geom: not geom.is_empty)] - - rows = [] + table = sdata.table[sdata.table.obs[_get_region_key(sdata)].isin([e])] - def assign_fill_and_outline_to_row( - shapes: list[GeoDataFrame], fill_c: list[Any], outline_c: list[Any], row: pd.Series, idx: int - ) -> None: - if len(shapes) > 1 and len(fill_c) == 1: - row["fill_c"] = fill_c - row["outline_c"] = outline_c - else: - row["fill_c"] = fill_c[idx] - row["outline_c"] = outline_c[idx] - - # Match colors to the geometry, potentially expanding the row in case of - # multipolygons - for idx, row in shapes_df.iterrows(): - geom = row["geometry"] - if geom.geom_type == "Polygon": - row = row.to_dict() - row["geometry"] = Polygon(geom.exterior.coords, closed=True) - assign_fill_and_outline_to_row(shapes, fill_c, outline_c, row, idx) - rows.append(row) - - elif geom.geom_type == "MultiPolygon": - mp = _make_patch_from_multipolygon(geom) - for _, m in enumerate(mp): - mp_copy = row.to_dict() - mp_copy["geometry"] = m - assign_fill_and_outline_to_row(shapes, fill_c, outline_c, mp_copy, idx) - rows.append(mp_copy) - - elif geom.geom_type == "Point": - row = row.to_dict() - row["geometry"] = Circle((geom.x, geom.y), radius=row["radius"]) - assign_fill_and_outline_to_row(shapes, fill_c, outline_c, row, idx) - rows.append(row) - - patches = pd.DataFrame(rows) - - return PatchCollection( - patches["geometry"].values.tolist(), - snap=False, - lw=render_params.outline_params.linewidth, - facecolor=patches["fill_c"], - edgecolor=None if all(outline is None for outline in outline_c) else outline_c, - **kwargs, + # get color vector (categorical or continuous) + color_source_vector, color_vector, _ = _set_color_source_vec( + sdata=sdata_filt, + element=sdata_filt.shapes[e], + element_name=e, + value_to_plot=render_params.color, + layer=render_params.layer, + groups=render_params.groups, + palette=render_params.palette, + na_color=render_params.cmap_params.na_color, + alpha=render_params.fill_alpha, ) - norm = copy(render_params.cmap_params.norm) - - if len(color_vector) == 0: - color_vector = [render_params.cmap_params.na_color] - - shapes = pd.concat(shapes, ignore_index=True) - shapes = gpd.GeoDataFrame(shapes, geometry="geometry") - _cax = _get_collection_shape( - shapes=shapes, - s=render_params.size, - c=color_vector, - rasterized=sc_settings._vector_friendly, - cmap=render_params.cmap_params.cmap, - norm=norm, - fill_alpha=render_params.fill_alpha, - outline_alpha=render_params.outline_alpha - # **kwargs, - ) - - 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 - - _ = _decorate_axs( - ax=ax, - cax=cax, - fig_params=fig_params, - adata=table, - value_to_plot=render_params.color, - color_source_vector=color_source_vector, - palette=palette, - alpha=render_params.fill_alpha, - na_color=render_params.cmap_params.na_color, - legend_fontsize=legend_params.legend_fontsize, - legend_fontweight=legend_params.legend_fontweight, - legend_loc=legend_params.legend_loc, - legend_fontoutline=legend_params.legend_fontoutline, - na_in_legend=legend_params.na_in_legend, - colorbar=legend_params.colorbar, - scalebar_dx=scalebar_params.scalebar_dx, - scalebar_units=scalebar_params.scalebar_units, - # scalebar_kwargs=scalebar_params.scalebar_kwargs, - ) - ax.set_aspect("equal") - ax.invert_yaxis() + # color_source_vector is None when the values aren't categorical + if color_source_vector is None and render_params.transfunc is not None: + color_vector = render_params.transfunc(color_vector) + + norm = copy(render_params.cmap_params.norm) + + if len(color_vector) == 0: + color_vector = [render_params.cmap_params.na_color] + + shapes = gpd.GeoDataFrame(shapes, geometry="geometry") + _cax = _get_collection_shape( + shapes=shapes, + s=render_params.size, + c=color_vector, + render_params=render_params, + rasterized=sc_settings._vector_friendly, + cmap=render_params.cmap_params.cmap, + norm=norm, + fill_alpha=render_params.fill_alpha, + outline_alpha=render_params.outline_alpha + # **kwargs, + ) + cax = ax.add_collection(_cax) -@dataclass -class PointsRenderParams: - """Points render parameters..""" + # 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 + ) - cmap_params: CmapParams - elements: str | Sequence[str] | None = None - color: str | None = None - groups: str | Sequence[str] | None = None - palette: ListedColormap | str | None = None - alpha: float = 1.0 - size: float = 1.0 - transfunc: Callable[[float], float] | None = None + # print(len(set(color_vector)) == 1) + # print(set(color_source_vector[0]) == to_hex(render_params.cmap_params.na_color)) + if not ( + len(set(color_vector)) == 1 and list(set(color_vector))[0] == to_hex(render_params.cmap_params.na_color) + ): + _ = _decorate_axs( + ax=ax, + cax=cax, + fig_params=fig_params, + adata=table, + value_to_plot=render_params.color, + color_source_vector=color_source_vector, + palette=palette, + alpha=render_params.fill_alpha, + na_color=render_params.cmap_params.na_color, + legend_fontsize=legend_params.legend_fontsize, + legend_fontweight=legend_params.legend_fontweight, + legend_loc=legend_params.legend_loc, + legend_fontoutline=legend_params.legend_fontoutline, + na_in_legend=legend_params.na_in_legend, + colorbar=legend_params.colorbar, + scalebar_dx=scalebar_params.scalebar_dx, + scalebar_units=scalebar_params.scalebar_units, + ) def _render_points( @@ -295,89 +167,80 @@ def _render_points( if elements is None: elements = list(sdata_filt.points.keys()) - points = [sdata.points[e] for e in elements] + for e in elements: + points = sdata.points[e] + coords = ["x", "y"] + if render_params.color is not None: + color = [render_params.color] if isinstance(render_params.color, str) else render_params.color + coords.extend(color) - coords = ["x", "y"] - if render_params.color is not None: - color = [render_params.color] if isinstance(render_params.color, str) else render_params.color - coords.extend(color) + point_df = points[coords].compute() - point_df = pd.concat([point[coords].compute() for point in points], axis=0) + # we construct an anndata to hack the plotting functions + adata = AnnData( + X=point_df[["x", "y"]].values, obs=point_df[coords].reset_index(), dtype=point_df[["x", "y"]].values.dtype + ) + if render_params.color is not None: + cols = sc.get.obs_df(adata, render_params.color) + # maybe set color based on type + if is_categorical_dtype(cols): + _maybe_set_colors( + source=adata, + target=adata, + key=render_params.color, + palette=render_params.palette, + ) + # print(p) + color_source_vector, color_vector, _ = _set_color_source_vec( + sdata=sdata_filt, + element=points, + element_name=e, + value_to_plot=render_params.color, + groups=render_params.groups, + palette=render_params.palette, + na_color=render_params.cmap_params.na_color, + alpha=render_params.alpha, + ) - # we construct an anndata to hack the plotting functions - adata = AnnData( - X=point_df[["x", "y"]].values, obs=point_df[coords].reset_index(), dtype=point_df[["x", "y"]].values.dtype - ) - if render_params.color is not None: - cols = sc.get.obs_df(adata, render_params.color) - # maybe set color based on type - if is_categorical_dtype(cols): - _maybe_set_colors( - source=adata, - target=adata, - key=render_params.color, + # color_source_vector is None when the values aren't categorical + if color_source_vector is None and render_params.transfunc is not None: + color_vector = render_params.transfunc(color_vector) + + norm = copy(render_params.cmap_params.norm) + _cax = ax.scatter( + adata[:, 0].X.flatten(), + adata[:, 1].X.flatten(), + s=render_params.size, + c=color_vector, + rasterized=sc_settings._vector_friendly, + cmap=render_params.cmap_params.cmap, + norm=norm, + alpha=render_params.alpha, + # **kwargs, + ) + cax = ax.add_collection(_cax) + if not ( + len(set(color_vector)) == 1 and list(set(color_vector))[0] == to_hex(render_params.cmap_params.na_color) + ): + _ = _decorate_axs( + ax=ax, + cax=cax, + fig_params=fig_params, + adata=adata, + value_to_plot=render_params.color, + color_source_vector=color_source_vector, palette=render_params.palette, + alpha=render_params.alpha, + na_color=render_params.cmap_params.na_color, + legend_fontsize=legend_params.legend_fontsize, + legend_fontweight=legend_params.legend_fontweight, + legend_loc=legend_params.legend_loc, + legend_fontoutline=legend_params.legend_fontoutline, + na_in_legend=legend_params.na_in_legend, + colorbar=legend_params.colorbar, + scalebar_dx=scalebar_params.scalebar_dx, + scalebar_units=scalebar_params.scalebar_units, ) - color_source_vector, color_vector, _ = _set_color_source_vec( - adata=adata, - value_to_plot=render_params.color, - groups=render_params.groups, - palette=render_params.palette, - na_color=render_params.cmap_params.na_color, - alpha=render_params.alpha, - ) - - # color_source_vector is None when the values aren't categorical - if color_source_vector is None and render_params.transfunc is not None: - color_vector = render_params.transfunc(color_vector) - - norm = copy(render_params.cmap_params.norm) - _cax = ax.scatter( - adata[:, 0].X.flatten(), - adata[:, 1].X.flatten(), - s=render_params.size, - c=color_vector, - rasterized=sc_settings._vector_friendly, - cmap=render_params.cmap_params.cmap, - norm=norm, - alpha=render_params.alpha, - # **kwargs, - ) - cax = ax.add_collection(_cax) - _ = _decorate_axs( - ax=ax, - cax=cax, - fig_params=fig_params, - adata=adata, - value_to_plot=render_params.color, - color_source_vector=color_source_vector, - palette=render_params.palette, - alpha=render_params.alpha, - na_color=render_params.cmap_params.na_color, - legend_fontsize=legend_params.legend_fontsize, - legend_fontweight=legend_params.legend_fontweight, - legend_loc=legend_params.legend_loc, - legend_fontoutline=legend_params.legend_fontoutline, - na_in_legend=legend_params.na_in_legend, - colorbar=legend_params.colorbar, - scalebar_dx=scalebar_params.scalebar_dx, - scalebar_units=scalebar_params.scalebar_units, - # scalebar_kwargs=scalebar_params.scalebar_kwargs, - ) - ax.set_aspect("equal") - ax.invert_yaxis() - - -@dataclass -class ImageRenderParams: - """Labels render parameters..""" - - cmap_params: list[CmapParams] | CmapParams - elements: str | Sequence[str] | None = None - channel: list[str] | list[int] | int | str | None = None - palette: ListedColormap | str | None = None - alpha: float = 1.0 - quantiles_for_norm: tuple[float | None, float | None] = (None, None) def _render_images( @@ -537,24 +400,6 @@ def _render_images( raise ValueError("If 'palette' is provided, 'cmap' must be None.") -@dataclass -class LabelsRenderParams: - """Labels render parameters..""" - - cmap_params: CmapParams - elements: str | Sequence[str] | None = None - color: str | None = None - groups: str | Sequence[str] | None = None - contour_px: int | None = None - outline: bool = False - alt_var: str | None = None - layer: str | None = None - palette: ListedColormap | str | None = None - outline_alpha: float = 1.0 - fill_alpha: float = 0.4 - transfunc: Callable[[float], float] | None = None - - def _render_labels( sdata: sd.SpatialData, render_params: LabelsRenderParams, @@ -597,9 +442,10 @@ def _render_labels( # get color vector (categorical or continuous) color_source_vector, color_vector, categorical = _set_color_source_vec( - adata=table, + sdata=sdata_filt, + element=sdata_filt.labels[label_key], + element_name=label_key, value_to_plot=render_params.color, - alt_var=render_params.alt_var, layer=render_params.layer, groups=render_params.groups, palette=render_params.palette, diff --git a/src/spatialdata_plot/pl/render_params.py b/src/spatialdata_plot/pl/render_params.py new file mode 100644 index 00000000..ac78eeb6 --- /dev/null +++ b/src/spatialdata_plot/pl/render_params.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +from collections.abc import Callable, Sequence +from dataclasses import dataclass +from typing import Literal + +from matplotlib.axes import Axes +from matplotlib.colors import Colormap, ListedColormap, Normalize +from matplotlib.figure import Figure + +_FontWeight = Literal["light", "normal", "medium", "semibold", "bold", "heavy", "black"] +_FontSize = Literal["xx-small", "x-small", "small", "medium", "large", "x-large", "xx-large"] + + +@dataclass +class CmapParams: + """Cmap params.""" + + cmap: Colormap + norm: Normalize + na_color: str | tuple[float, ...] = (0.0, 0.0, 0.0, 0.0) + + +@dataclass +class FigParams: + """Figure params.""" + + fig: Figure + ax: Axes + num_panels: int + axs: Sequence[Axes] | None = None + title: str | Sequence[str] | None = None + ax_labels: Sequence[str] | None = None + frameon: bool | None = None + + +@dataclass +class OutlineParams: + """Cmap params.""" + + outline: bool + outline_color: str | list[float] + linewidth: float + + +@dataclass +class LegendParams: + """Legend params.""" + + legend_fontsize: int | float | _FontSize | None = None + legend_fontweight: int | _FontWeight = "bold" + legend_loc: str | None = "right margin" + legend_fontoutline: int | None = None + na_in_legend: bool = True + colorbar: bool = True + + +@dataclass +class ScalebarParams: + """Scalebar params.""" + + scalebar_dx: Sequence[float] | None = None + scalebar_units: Sequence[str] | None = None + + +@dataclass +class ShapesRenderParams: + """Labels render parameters..""" + + cmap_params: CmapParams + outline_params: OutlineParams + elements: str | Sequence[str] | None = None + color: str | None = None + groups: str | Sequence[str] | None = None + contour_px: int | None = None + layer: str | None = None + palette: ListedColormap | str | None = None + outline_alpha: float = 1.0 + fill_alpha: float = 0.3 + size: float = 1.0 + transfunc: Callable[[float], float] | None = None + + +@dataclass +class PointsRenderParams: + """Points render parameters..""" + + cmap_params: CmapParams + elements: str | Sequence[str] | None = None + color: str | None = None + groups: str | Sequence[str] | None = None + palette: ListedColormap | str | None = None + alpha: float = 1.0 + size: float = 1.0 + transfunc: Callable[[float], float] | None = None + + +@dataclass +class ImageRenderParams: + """Labels render parameters..""" + + cmap_params: list[CmapParams] | CmapParams + elements: str | Sequence[str] | None = None + channel: list[str] | list[int] | int | str | None = None + palette: ListedColormap | str | None = None + alpha: float = 1.0 + quantiles_for_norm: tuple[float | None, float | None] = (None, None) + + +@dataclass +class LabelsRenderParams: + """Labels render parameters..""" + + cmap_params: CmapParams + elements: str | Sequence[str] | None = None + color: str | None = None + groups: str | Sequence[str] | None = None + contour_px: int | None = None + outline: bool = False + layer: str | None = None + palette: ListedColormap | str | None = None + outline_alpha: float = 1.0 + fill_alpha: float = 0.4 + transfunc: Callable[[float], float] | None = None diff --git a/src/spatialdata_plot/pl/utils.py b/src/spatialdata_plot/pl/utils.py index ee9e6541..32324597 100644 --- a/src/spatialdata_plot/pl/utils.py +++ b/src/spatialdata_plot/pl/utils.py @@ -3,7 +3,6 @@ import os from collections.abc import Iterable, Mapping, Sequence from copy import copy -from dataclasses import dataclass from functools import partial from pathlib import Path from types import MappingProxyType @@ -11,6 +10,7 @@ import matplotlib import matplotlib.patches as mpatches +import matplotlib.patches as mplp import matplotlib.path as mpath import matplotlib.pyplot as plt import multiscale_spatial_image as msi @@ -22,10 +22,19 @@ import xarray as xr from anndata import AnnData from cycler import Cycler, cycler +from geopandas import GeoDataFrame from matplotlib import colors, patheffects, rcParams from matplotlib.axes import Axes from matplotlib.collections import PatchCollection -from matplotlib.colors import Colormap, LinearSegmentedColormap, ListedColormap, Normalize, TwoSlopeNorm, to_rgba +from matplotlib.colors import ( + ColorConverter, + Colormap, + LinearSegmentedColormap, + ListedColormap, + Normalize, + TwoSlopeNorm, + to_rgba, +) from matplotlib.figure import Figure from matplotlib.gridspec import GridSpec from matplotlib_scalebar.scalebar import ScaleBar @@ -40,40 +49,26 @@ from skimage.segmentation import find_boundaries from skimage.util import map_array from spatialdata import transform +from spatialdata._core.query.relational_query import _locate_value, get_values from spatialdata._logging import logger as logging from spatialdata._types import ArrayLike -from spatialdata.models import Image2DModel, Labels2DModel +from spatialdata.models import Image2DModel, Labels2DModel, SpatialElement from spatialdata.transformations import get_transformation +from spatialdata_plot.pl.render_params import ( + CmapParams, + FigParams, + OutlineParams, + ScalebarParams, + ShapesRenderParams, + _FontSize, + _FontWeight, +) from spatialdata_plot.pp.utils import _get_coordinate_system_mapping -_FontWeight = Literal["light", "normal", "medium", "semibold", "bold", "heavy", "black"] -_FontSize = Literal["xx-small", "x-small", "small", "medium", "large", "x-large", "xx-large"] - to_hex = partial(colors.to_hex, keep_alpha=True) -@dataclass -class FigParams: - """Figure params.""" - - fig: Figure - ax: Axes - num_panels: int - axs: Sequence[Axes] | None = None - title: str | Sequence[str] | None = None - ax_labels: Sequence[str] | None = None - frameon: bool | None = None - - -@dataclass -class ScalebarParams: - """Scalebar params.""" - - scalebar_dx: Sequence[float] | None = None - scalebar_units: Sequence[str] | None = None - - def _prepare_params_plot( # this param is inferred when `pl.show`` is called num_panels: int, @@ -174,6 +169,108 @@ def _get_cs_contents(sdata: sd.SpatialData) -> pd.DataFrame: return cs_contents +def _get_collection_shape( + shapes: list[GeoDataFrame], + c: Any, + s: float, + norm: Any, + render_params: ShapesRenderParams, + fill_alpha: None | float = None, + outline_alpha: None | float = None, + **kwargs: Any, +) -> PatchCollection: + """ + Get a PatchCollection for rendering given geometries with specified colors and outlines. + + Args: + - shapes (list[GeoDataFrame]): List of geometrical shapes. + - c: Color parameter. + - s (float): Size of the shape. + - norm: Normalization for the color map. + - fill_alpha (float, optional): Opacity for the fill color. + - outline_alpha (float, optional): Opacity for the outline. + - **kwargs: Additional keyword arguments. + + Returns + ------- + - PatchCollection: Collection of patches for rendering. + """ + cmap = kwargs["cmap"] + + try: + # fails when numeric + fill_c = ColorConverter().to_rgba_array(c) + except ValueError: + if norm is None: + c = cmap(c) + else: + norm = colors.Normalize(vmin=min(c), vmax=max(c)) + c = cmap(norm(c)) + + fill_c = ColorConverter().to_rgba_array(c) + fill_c[..., -1] = render_params.fill_alpha + + if render_params.outline_params.outline: + outline_c = ColorConverter().to_rgba_array(render_params.outline_params.outline_color) + outline_c[..., -1] = render_params.outline_alpha + outline_c = outline_c.tolist() + else: + outline_c = [None] + outline_c = outline_c * fill_c.shape[0] + + shapes_df = pd.DataFrame(shapes, copy=True) + + # remove empty points/polygons + shapes_df = shapes_df[shapes_df["geometry"].apply(lambda geom: not geom.is_empty)] + + rows = [] + + def assign_fill_and_outline_to_row( + shapes: list[GeoDataFrame], fill_c: list[Any], outline_c: list[Any], row: pd.Series, idx: int + ) -> None: + if len(shapes) > 1 and len(fill_c) == 1: + row["fill_c"] = fill_c + row["outline_c"] = outline_c + else: + row["fill_c"] = fill_c[idx] + row["outline_c"] = outline_c[idx] + + # Match colors to the geometry, potentially expanding the row in case of + # multipolygons + for idx, row in shapes_df.iterrows(): + geom = row["geometry"] + if geom.geom_type == "Polygon": + row = row.to_dict() + row["geometry"] = mplp.Polygon(geom.exterior.coords, closed=True) + assign_fill_and_outline_to_row(shapes, fill_c, outline_c, row, idx) + rows.append(row) + + elif geom.geom_type == "MultiPolygon": + mp = _make_patch_from_multipolygon(geom) + for _, m in enumerate(mp): + mp_copy = row.to_dict() + mp_copy["geometry"] = m + assign_fill_and_outline_to_row(shapes, fill_c, outline_c, mp_copy, idx) + rows.append(mp_copy) + + elif geom.geom_type == "Point": + row = row.to_dict() + row["geometry"] = mplp.Circle((geom.x, geom.y), radius=row["radius"]) + assign_fill_and_outline_to_row(shapes, fill_c, outline_c, row, idx) + rows.append(row) + + patches = pd.DataFrame(rows) + + return PatchCollection( + patches["geometry"].values.tolist(), + snap=False, + lw=render_params.outline_params.linewidth, + facecolor=patches["fill_c"], + edgecolor=None if all(outline is None for outline in outline_c) else outline_c, + **kwargs, + ) + + def _get_extent( sdata: sd.SpatialData, coordinate_systems: Sequence[str] | str | None = None, @@ -456,15 +553,6 @@ def _get_scalebar( return _scalebar_dx, _scalebar_units -@dataclass -class CmapParams: - """Cmap params.""" - - cmap: Colormap - norm: Normalize - na_color: str | tuple[float, ...] = (0.0, 0.0, 0.0, 0.0) - - def _prepare_cmap_norm( cmap: Colormap | str | None = None, norm: Normalize | Sequence[Normalize] | None = None, @@ -487,15 +575,6 @@ def _prepare_cmap_norm( return CmapParams(cmap, norm, na_color) -@dataclass -class OutlineParams: - """Cmap params.""" - - outline: bool - outline_color: str | list[float] - linewidth: float - - def _set_outline( size: float, outline: bool = False, @@ -714,10 +793,10 @@ def _get_colors_for_categorical_obs( def _set_color_source_vec( - adata: AnnData, + sdata: sd.SpatialData, + element: SpatialElement | None, value_to_plot: str | None, - use_raw: bool | None = None, - alt_var: str | None = None, + element_name: list[str] | str | None = None, layer: str | None = None, groups: Sequence[str] | str | None = None, palette: ListedColormap | str | None = None, @@ -725,39 +804,54 @@ def _set_color_source_vec( alpha: float = 1.0, ) -> tuple[ArrayLike | pd.Series | None, ArrayLike, bool]: if value_to_plot is None: - color = np.full(adata.n_obs, to_hex(na_color)) + color = np.full(len(element), to_hex(na_color)) # type: ignore[arg-type] return color, color, False - if alt_var is not None and value_to_plot not in adata.obs and value_to_plot not in adata.var_names: - value_to_plot = adata.var_names[adata.var[alt_var] == value_to_plot][0] - if use_raw and value_to_plot not in adata.obs: - color_source_vector = adata.raw.obs_vector(value_to_plot) - else: - color_source_vector = adata.obs_vector(value_to_plot, layer=layer) + # Figure out where to get the color from + origins = _locate_value(value_key=value_to_plot, sdata=sdata, element_name=element_name) + if len(origins) > 1: + raise ValueError( + f"Color key '{value_to_plot}' for element '{element_name}' been found in multiple locations: {origins}." + ) - if not is_categorical_dtype(color_source_vector): - return None, color_source_vector, False + if len(origins) == 1: + vals = get_values(value_key=value_to_plot, sdata=sdata, element_name=element_name) + color_source_vector = vals[value_to_plot] - color_source_vector = pd.Categorical(color_source_vector) # convert, e.g., `pd.Series` - categories = color_source_vector.categories + # if all([isinstance(x, str) for x in color_source_vector]): + # raise TypeError( + # f"Color key '{value_to_plot}' for element '{element_name}' has string values, " + # f"but should be numerical or categorical." + # ) - if groups is not None: - color_source_vector = color_source_vector.remove_categories(categories.difference(groups)) + # numerical case, return early + if not is_categorical_dtype(color_source_vector): + return None, color_source_vector, False - color_map = dict(zip(categories, _get_colors_for_categorical_obs(categories))) - # color_map = _get_palette( - # adata=adata, cluster_key=value_to_plot, categories=categories, palette=palette, alpha=alpha - # ) - if color_map is None: - raise ValueError("Unable to create color palette.") + color_source_vector = pd.Categorical(color_source_vector) # convert, e.g., `pd.Series` + categories = color_source_vector.categories - # do not rename categories, as colors need not be unique - color_vector = color_source_vector.map(color_map) - if color_vector.isna().any(): - color_vector = color_vector.add_categories([to_hex(na_color)]) - color_vector = color_vector.fillna(to_hex(na_color)) + if groups is not None: + color_source_vector = color_source_vector.remove_categories(categories.difference(groups)) - return color_source_vector, color_vector, True + color_map = dict(zip(categories, _get_colors_for_categorical_obs(categories))) + # color_map = _get_palette( + # adata=adata, cluster_key=value_to_plot, categories=categories, palette=palette, alpha=alpha + # ) + if color_map is None: + raise ValueError("Unable to create color palette.") + + # do not rename categories, as colors need not be unique + color_vector = color_source_vector.map(color_map) + if color_vector.isna().any(): + color_vector = color_vector.add_categories([to_hex(na_color)]) + color_vector = color_vector.fillna(to_hex(na_color)) + + return color_source_vector, color_vector, True + + logging.warning(f"Color key '{value_to_plot}' for element '{element_name}' not been found, using default colors.") + color = np.full(sdata.table.n_obs, to_hex(na_color)) + return color, color, False def _map_color_seg( @@ -871,18 +965,6 @@ def _maybe_set_colors( add_colors_for_categorical_sample_annotation(target, key=key, force_update_colors=True, palette=palette) -@dataclass -class LegendParams: - """Legend params.""" - - legend_fontsize: int | float | _FontSize | None = None - legend_fontweight: int | _FontWeight = "bold" - legend_loc: str | None = "right margin" - legend_fontoutline: int | None = None - na_in_legend: bool = True - colorbar: bool = True - - def _decorate_axs( ax: Axes, cax: PatchCollection, diff --git a/tests/_images/Images_can_normalize_image.png b/tests/_images/Images_can_normalize_image.png new file mode 100644 index 00000000..d83f4034 Binary files /dev/null and b/tests/_images/Images_can_normalize_image.png differ diff --git a/tests/_images/Shapes_can_color_from_geodataframe.png b/tests/_images/Shapes_can_color_from_geodataframe.png new file mode 100644 index 00000000..c9718063 Binary files /dev/null and b/tests/_images/Shapes_can_color_from_geodataframe.png differ diff --git a/tests/pl/test_get_extent.py b/tests/pl/test_get_extent.py index ad82453b..58b664d2 100644 --- a/tests/pl/test_get_extent.py +++ b/tests/pl/test_get_extent.py @@ -42,15 +42,3 @@ def test_plot_extent_of_img_is_correct_after_spatial_query(self, sdata_blobs: Sp axes=["x", "y"], min_coordinate=[100, 100], max_coordinate=[400, 400], target_coordinate_system="global" ) cropped_blobs.pl.render_images().pl.show() - - def test_plot_extent_of_polygons_is_correct_after_spatial_query(self, sdata_blobs: SpatialData): - cropped_blobs = sdata_blobs.pp.get_elements(["blobs_polygons"]).query.bounding_box( - axes=["x", "y"], min_coordinate=[100, 100], max_coordinate=[400, 400], target_coordinate_system="global" - ) - cropped_blobs.pl.render_shapes().pl.show() - - def test_plot_extent_of_polygons_on_img_is_correct_after_spatial_query(self, sdata_blobs: SpatialData): - cropped_blobs = sdata_blobs.pp.get_elements(["blobs_image", "blobs_polygons"]).query.bounding_box( - axes=["x", "y"], min_coordinate=[100, 100], max_coordinate=[400, 400], target_coordinate_system="global" - ) - cropped_blobs.pl.render_images().pl.render_shapes().pl.show() diff --git a/tests/pl/test_render_images.py b/tests/pl/test_render_images.py index 20acfbde..6189e189 100644 --- a/tests/pl/test_render_images.py +++ b/tests/pl/test_render_images.py @@ -46,3 +46,6 @@ def test_plot_can_pass_cmap_to_each_channel(self, sdata_blobs: SpatialData): sdata_blobs.pl.render_images( elements="blobs_image", channel=[0, 1, 2], cmap=["Reds", "Greens", "Blues"] ).pl.show() + + def test_plot_can_normalize_image(self, sdata_blobs: SpatialData): + sdata_blobs.pl.render_images(elements="blobs_image", quantiles_for_norm=(5, 90)).pl.show() diff --git a/tests/pl/test_render_shapes.py b/tests/pl/test_render_shapes.py index 5f9ea0c1..e183cea7 100644 --- a/tests/pl/test_render_shapes.py +++ b/tests/pl/test_render_shapes.py @@ -85,7 +85,15 @@ def _make_multi(): sdata = SpatialData(shapes={"p": _make_multi()}) adata = anndata.AnnData(pd.DataFrame({"p": ["hole", "overlap", "square", "circle"]})) adata.obs.loc[:, "region"] = "p" - adata.obs.loc[:, "val"] = [1, 2, 3, 4] + adata.obs.loc[:, "val"] = [0, 1, 2, 3] table = TableModel.parse(adata, region="p", region_key="region", instance_key="val") sdata.table = table sdata.pl.render_shapes(color="val", outline=True, fill_alpha=0.3).pl.show() + + def test_plot_can_color_from_geodataframe(self, sdata_blobs: SpatialData): + blob = sdata_blobs + blob.shapes["blobs_polygons"]["value"] = [1, 10, 1, 20, 1] + blob.pl.render_shapes( + elements="blobs_polygons", + color="value", + ).pl.show()