From 5ec23c2ca6fc928c388d5c4684407cd0c4b45757 Mon Sep 17 00:00:00 2001 From: Tim Treis Date: Mon, 4 Sep 2023 10:57:28 +0200 Subject: [PATCH 1/2] mvp --- src/spatialdata_plot/pl/basic.py | 1 + src/spatialdata_plot/pl/render.py | 13 ++++++++++++- 2 files changed, 13 insertions(+), 1 deletion(-) diff --git a/src/spatialdata_plot/pl/basic.py b/src/spatialdata_plot/pl/basic.py index f09425b5..517967aa 100644 --- a/src/spatialdata_plot/pl/basic.py +++ b/src/spatialdata_plot/pl/basic.py @@ -203,6 +203,7 @@ def render_shapes( ------- None """ + sdata = self._copy() sdata = _verify_plotting_tree(sdata) n_steps = len(sdata.plotting_tree.keys()) diff --git a/src/spatialdata_plot/pl/render.py b/src/spatialdata_plot/pl/render.py index 6e264532..97ba337b 100644 --- a/src/spatialdata_plot/pl/render.py +++ b/src/spatialdata_plot/pl/render.py @@ -8,6 +8,7 @@ import matplotlib import numpy as np import pandas as pd +import dask import scanpy as sc import spatial_image import spatialdata as sd @@ -18,6 +19,7 @@ from spatialdata.models import ( Image2DModel, Labels2DModel, + PointsModel, ) from spatialdata_plot._logging import logger @@ -155,6 +157,10 @@ def _render_points( scalebar_params: ScalebarParams, legend_params: LegendParams, ) -> 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( @@ -173,6 +179,12 @@ def _render_points( if render_params.color is not None: 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) + points = points[points[color].isin(render_params.groups).values] + + points = dask.dataframe.from_pandas(points, npartitions=1) + sdata_filt.points[e] = PointsModel.parse(points) point_df = points[coords].compute() @@ -190,7 +202,6 @@ def _render_points( key=render_params.color, palette=render_params.palette, ) - # print(p) color_source_vector, color_vector, _ = _set_color_source_vec( sdata=sdata_filt, element=points, From 0313cc39259b44afb83ff7293a08c3997d35036f Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Mon, 4 Sep 2023 08:58:10 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- src/spatialdata_plot/pl/basic.py | 1 - src/spatialdata_plot/pl/render.py | 4 ++-- 2 files changed, 2 insertions(+), 3 deletions(-) diff --git a/src/spatialdata_plot/pl/basic.py b/src/spatialdata_plot/pl/basic.py index 517967aa..f09425b5 100644 --- a/src/spatialdata_plot/pl/basic.py +++ b/src/spatialdata_plot/pl/basic.py @@ -203,7 +203,6 @@ def render_shapes( ------- None """ - sdata = self._copy() sdata = _verify_plotting_tree(sdata) n_steps = len(sdata.plotting_tree.keys()) diff --git a/src/spatialdata_plot/pl/render.py b/src/spatialdata_plot/pl/render.py index 97ba337b..e082b287 100644 --- a/src/spatialdata_plot/pl/render.py +++ b/src/spatialdata_plot/pl/render.py @@ -4,11 +4,11 @@ from copy import copy from typing import Union +import dask import geopandas as gpd import matplotlib import numpy as np import pandas as pd -import dask import scanpy as sc import spatial_image import spatialdata as sd @@ -182,7 +182,7 @@ def _render_points( points = points[coords].compute() points[color[0]].cat.set_categories(render_params.groups, inplace=True) points = points[points[color].isin(render_params.groups).values] - + points = dask.dataframe.from_pandas(points, npartitions=1) sdata_filt.points[e] = PointsModel.parse(points)