diff --git a/.gitignore b/.gitignore index 84ad50c4..c8424e22 100644 --- a/.gitignore +++ b/.gitignore @@ -8,6 +8,7 @@ buck-out/ __pycache__/ .mypy_cache/ .ruff_cache/ +/node_modules # Distribution / packaging /build/ diff --git a/src/spatialdata_plot/pl/basic.py b/src/spatialdata_plot/pl/basic.py index 6aea98db..5ff6e77e 100644 --- a/src/spatialdata_plot/pl/basic.py +++ b/src/spatialdata_plot/pl/basic.py @@ -272,7 +272,7 @@ def render_shapes( def render_points( self, elements: list[str] | str | None = None, - color: list[str] | str | None = None, + color: list[str | None] | str | None = None, alpha: float | int = 1.0, groups: list[list[str | None]] | list[str] | str | None = None, palette: list[list[str | None]] | list[str] | str | None = None, @@ -475,7 +475,7 @@ def render_images( def render_labels( self, elements: list[str] | str | None = None, - color: list[str] | str | None = None, + color: list[str | None] | str | None = None, groups: list[list[str | None]] | list[str] | str | None = None, contour_px: int = 3, outline: bool = False, @@ -562,6 +562,7 @@ def render_labels( na_color=na_color, # type: ignore[arg-type] **kwargs, ) + sdata.plotting_tree[f"{n_steps+1}_render_labels"] = LabelsRenderParams( elements=params_dict["elements"], color=params_dict["color"], diff --git a/src/spatialdata_plot/pl/render.py b/src/spatialdata_plot/pl/render.py index 736d3aa5..b82fe19c 100644 --- a/src/spatialdata_plot/pl/render.py +++ b/src/spatialdata_plot/pl/render.py @@ -44,6 +44,8 @@ _multiscale_to_spatial_image, _normalize, _rasterize_if_necessary, + _return_list_list_str_none, + _return_list_str_none, _set_color_source_vec, to_hex, ) @@ -62,7 +64,12 @@ def _render_shapes( ) -> None: elements = render_params.elements element_table_mapping = render_params.element_table_mapping + cols_for_color = _return_list_str_none(render_params.col_for_color) + colors = _return_list_str_none(render_params.color) + groups = _return_list_list_str_none(render_params.groups) + palettes = _return_list_list_str_none(render_params.palette) + assert isinstance(element_table_mapping, dict) sdata_filt = sdata.filter_by_coordinate_system( coordinate_system=coordinate_system, filter_tables=any(value is not None for value in element_table_mapping.values()), @@ -72,7 +79,7 @@ def _render_shapes( elements = list(sdata_filt.shapes.keys()) for index, e in enumerate(elements): - col_for_color = render_params.col_for_color[index] + col_for_color = cols_for_color[index] shapes = sdata.shapes[e] table_name = element_table_mapping.get(e) @@ -104,13 +111,13 @@ def _render_shapes( element_index=index, element_name=e, value_to_plot=col_for_color, - groups=render_params.groups[index] if render_params.groups[index][0] is not None else None, + groups=groups[index] if groups[index][0] is not None else None, palette=( - render_params.palette[index] if render_params.palette is not None else None + palettes[index] if palettes is not None else None ), # and render_params.palette[index][0] is not None - na_color=render_params.color[index] or render_params.cmap_params.na_color, + na_color=colors[index] or render_params.cmap_params.na_color, cmap_params=render_params.cmap_params, - table_name=table_name, + table_name=cast(str, table_name), ) values_are_categorical = color_source_vector is not None @@ -126,12 +133,8 @@ def _render_shapes( # filter by `groups` - if ( - isinstance(render_params.groups, list) - and render_params.groups[index][0] is not None - and color_source_vector is not None - ): - mask = color_source_vector.isin(render_params.groups[index]) + if isinstance(groups, list) and groups[index][0] is not None and color_source_vector is not None: + mask = color_source_vector.isin(groups[index]) shapes = shapes[mask] shapes = shapes.reset_index() color_source_vector = color_source_vector[mask] @@ -177,11 +180,11 @@ def _render_shapes( len(set(color_vector)) == 1 and list(set(color_vector))[0] == to_hex(render_params.cmap_params.na_color) ): # necessary in case different shapes elements are annotated with one table - if color_source_vector is not None and render_params.col_for_color[index] is not None: + if color_source_vector is not None and col_for_color is not None: color_source_vector = color_source_vector.remove_unused_categories() # False if user specified color-like with 'color' parameter - colorbar = False if render_params.col_for_color[index] is None else legend_params.colorbar + colorbar = False if cols_for_color[index] is None else legend_params.colorbar _ = _decorate_axs( ax=ax, @@ -215,6 +218,12 @@ def _render_points( ) -> None: elements = render_params.elements element_table_mapping = render_params.element_table_mapping + cols_for_color = _return_list_str_none(render_params.col_for_color) + groups = _return_list_list_str_none(render_params.groups) + palettes = _return_list_list_str_none(render_params.palette) + colors = _return_list_str_none(render_params.color) + # Purely for mypy + assert isinstance(element_table_mapping, dict) sdata_filt = sdata.filter_by_coordinate_system( coordinate_system=coordinate_system, @@ -226,11 +235,11 @@ def _render_points( for index, e in enumerate(elements): points = sdata.points[e] - col_for_color = render_params.col_for_color[index] + col_for_color = cols_for_color[index] table_name = element_table_mapping.get(e) coords = ["x", "y"] - # if col_for_color is not None: + if ( col_for_color is not None and col_for_color not in points.columns @@ -257,8 +266,8 @@ def _render_points( coords += [col_for_color] points = points[coords].compute() - if render_params.groups[index][0] is not None and col_for_color is not None: - points = points[points[col_for_color].isin(render_params.groups[index])] + if groups[index][0] is not None and col_for_color is not None: + points = points[points[col_for_color].isin(groups[index])] # we construct an anndata to hack the plotting functions if table_name is None: @@ -285,14 +294,12 @@ def _render_points( source=adata, target=adata, key=col_for_color, - palette=render_params.palette[index] if render_params.palette[index][0] is not None else None, + palette=palettes[index] if palettes[index][0] is not None else None, ) # when user specified a single color, we overwrite na with it default_color = ( - render_params.color[index] - if col_for_color is None and render_params.color[index] is not None - else render_params.cmap_params.na_color + colors[index] if col_for_color is None and colors[index] is not None else render_params.cmap_params.na_color ) color_source_vector, color_vector, _ = _set_color_source_vec( @@ -300,9 +307,9 @@ def _render_points( element=points, element_index=index, element_name=e, - value_to_plot=render_params.col_for_color[index], - groups=render_params.groups[index] if render_params.groups[index][0] is not None else None, - palette=render_params.palette[index] if render_params.palette[index][0] is not None else None, + value_to_plot=col_for_color, + groups=groups[index] if groups[index][0] is not None else None, + palette=palettes[index] if palettes[index][0] is not None else None, na_color=default_color, cmap_params=render_params.cmap_params, table_name=cast(str, table_name), @@ -327,7 +334,6 @@ def _render_points( norm=norm, alpha=render_params.alpha, transform=trans, - # **kwargs, ) cax = ax.add_collection(_cax) @@ -342,7 +348,7 @@ def _render_points( cax=cax, fig_params=fig_params, adata=adata, - value_to_plot=render_params.col_for_color, + value_to_plot=col_for_color, color_source_vector=color_source_vector, palette=palette, alpha=render_params.alpha, @@ -369,6 +375,7 @@ def _render_images( rasterize: bool, ) -> None: elements = render_params.elements + palettes = _return_list_list_str_none(render_params.palette) sdata_filt = sdata.filter_by_coordinate_system( coordinate_system=coordinate_system, @@ -445,11 +452,11 @@ def _render_images( if render_params.cmap_params.norm is not None: # type: ignore[attr-defined] layer = render_params.cmap_params.norm(layer) # type: ignore[attr-defined] - if isinstance(render_params.palette, list): - if render_params.palette[i][0] is None: + if isinstance(palettes, list): + if palettes[i][0] is None: cmap = render_params.cmap_params.cmap # type: ignore[attr-defined] else: - cmap = _get_linear_colormap(render_params.palette[i], "k")[0] # type: ignore[arg-type] + cmap = _get_linear_colormap(palettes[i], "k")[0] # type: ignore[arg-type] # Overwrite alpha in cmap: https://stackoverflow.com/a/10127675 cmap._init() @@ -483,12 +490,8 @@ def _render_images( layers[c] = render_params.cmap_params[ch_index].norm(layers[c]) # 2A) Image has 3 channels, no palette info, and no/only one cmap was given - if isinstance(render_params.palette, list): - if ( - n_channels == 3 - and render_params.palette[i][0] is None - and not isinstance(render_params.cmap_params, list) - ): + if isinstance(palettes, list): + if n_channels == 3 and palettes[i][0] is None and not isinstance(render_params.cmap_params, list): if render_params.cmap_params.is_default: # -> use RGB stacked = np.stack([layers[c] for c in channels], axis=-1) else: # -> use given cmap for each channel @@ -516,7 +519,7 @@ def _render_images( im.set_transform(trans_data) # 2B) Image has n channels, no palette/cmap info -> sample n categorical colors - elif render_params.palette[i][0] is None and not got_multiple_cmaps: + elif palettes[i][0] is None and not got_multiple_cmaps: # overwrite if n_channels == 2 for intuitive result if n_channels == 2: seed_colors = ["#ff0000ff", "#00ff00ff"] @@ -538,11 +541,11 @@ def _render_images( im.set_transform(trans_data) # 2C) Image has n channels and palette info - elif render_params.palette[i][0] is not None and not got_multiple_cmaps: - if len(render_params.palette[i]) != n_channels: + elif palettes[i][0] is not None and not got_multiple_cmaps: + if len(palettes[i]) != n_channels: raise ValueError("If 'palette' is provided, its length must match the number of channels.") - channel_cmaps = [_get_linear_colormap([c], "k")[0] for c in render_params.palette[i]] + channel_cmaps = [_get_linear_colormap([c], "k")[0] for c in palettes[i] if isinstance(c, str)] # Apply cmaps to each channel and add up colored = np.stack([channel_cmaps[i](layers[c]) for i, c in enumerate(channels)], 0).sum(0) @@ -556,7 +559,7 @@ def _render_images( ) im.set_transform(trans_data) - elif render_params.palette[i][0] is None and got_multiple_cmaps: + elif palettes[i][0] is None and got_multiple_cmaps: channel_cmaps = [cp.cmap for cp in render_params.cmap_params] # type: ignore[union-attr] # Apply cmaps to each channel, add up and normalize to [0, 1] @@ -574,7 +577,7 @@ def _render_images( ) im.set_transform(trans_data) - elif render_params.palette[i][0] is not None and got_multiple_cmaps: + elif palettes[i][0] is not None and got_multiple_cmaps: raise ValueError("If 'palette' is provided, 'cmap' must be None.") @@ -590,6 +593,9 @@ def _render_labels( ) -> None: elements = render_params.elements element_table_mapping = cast(dict[str, str], render_params.element_table_mapping) + palettes = _return_list_list_str_none(render_params.palette) + colors = _return_list_str_none(render_params.color) + groups = _return_list_list_str_none(render_params.groups) if render_params.outline is False: render_params.outline_alpha = 0 @@ -606,7 +612,7 @@ def _render_labels( label = sdata_filt.labels[e] extent = get_extent(label, coordinate_system=coordinate_system) scale = render_params.scale[i] if isinstance(render_params.scale, list) else render_params.scale - color = render_params.color[i] + color = colors[i] # get best scale out of multiscale label if isinstance(label, MultiscaleSpatialImage): @@ -651,8 +657,8 @@ def _render_labels( element_index=i, element_name=e, value_to_plot=color, - groups=render_params.groups[i], - palette=render_params.palette[i], + groups=groups[i], # if isinstance(groups, list) else None, + palette=palettes[i], na_color=render_params.cmap_params.na_color, cmap_params=render_params.cmap_params, table_name=cast(str, table_name), @@ -733,7 +739,7 @@ def _render_labels( adata=table, value_to_plot=color, color_source_vector=color_source_vector, - palette=render_params.palette[i], + palette=palettes[i], alpha=render_params.fill_alpha, na_color=render_params.cmap_params.na_color, legend_fontsize=legend_params.legend_fontsize, diff --git a/src/spatialdata_plot/pl/render_params.py b/src/spatialdata_plot/pl/render_params.py index 6b838f1e..bcfa92e7 100644 --- a/src/spatialdata_plot/pl/render_params.py +++ b/src/spatialdata_plot/pl/render_params.py @@ -71,16 +71,16 @@ class ShapesRenderParams: cmap_params: CmapParams outline_params: OutlineParams elements: str | Sequence[str] | None = None - color: str | None = None + color: list[str | None] | str | None = None col_for_color: str | None = None - groups: str | Sequence[str] | None = None + groups: str | list[list[str | None]] | list[str | None] | None = None contour_px: int | None = None - palette: ListedColormap | str | None = None + palette: ListedColormap | list[str | None] | None = None outline_alpha: float = 1.0 fill_alpha: float = 0.3 scale: float = 1.0 transfunc: Callable[[float], float] | None = None - element_table_mapping: dict[str, set[str] | str] | str | list[str] | None = None + element_table_mapping: dict[str, set[str | None] | str | None] | str | list[str] | None = None @dataclass @@ -88,15 +88,15 @@ class PointsRenderParams: """Points render parameters..""" cmap_params: CmapParams - elements: str | Sequence[str] | None = None - color: str | None = None - col_for_color: str | None = None - groups: str | Sequence[str] | None = None - palette: ListedColormap | str | None = None + elements: str | list[str] | None = None + color: list[str | None] | str | None = None + col_for_color: list[str | None] | str | None = None + groups: str | list[list[str | None]] | list[str | None] | None = None + palette: ListedColormap | list[list[str | None]] | list[str | None] | None = None alpha: float = 1.0 size: float = 1.0 transfunc: Callable[[float], float] | None = None - element_table_mapping: dict[str, set[str] | str] | str | list[str] | None = None + element_table_mapping: dict[str, set[str | None] | str | None] | str | list[str] | None = None @dataclass @@ -106,7 +106,7 @@ class ImageRenderParams: 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 + palette: ListedColormap | list[str | None] | None = None alpha: float = 1.0 quantiles_for_norm: tuple[float | None, float | None] = (None, None) scale: str | list[str] | None = None @@ -119,12 +119,12 @@ class LabelsRenderParams: cmap_params: CmapParams elements: str | Sequence[str] | None = None color: list[str | None] | str | None = None - groups: str | Sequence[str] | None = None + groups: str | list[list[str | None]] | list[str | None] | None = None contour_px: int | None = None outline: bool = False - palette: ListedColormap | str | None = None + palette: ListedColormap | list[list[str | None]] | list[str | None] | None = None outline_alpha: float = 1.0 fill_alpha: float = 0.4 transfunc: Callable[[float], float] | None = None scale: str | list[str] | None = None - element_table_mapping: dict[str, set[str] | str] | str | list[str] | None = None + element_table_mapping: dict[str, set[str | None] | str | None] | str | list[str] | None = None diff --git a/src/spatialdata_plot/pl/utils.py b/src/spatialdata_plot/pl/utils.py index 62d7f2d5..5da2b173 100644 --- a/src/spatialdata_plot/pl/utils.py +++ b/src/spatialdata_plot/pl/utils.py @@ -8,7 +8,7 @@ from functools import partial from pathlib import Path from types import MappingProxyType -from typing import Any, Literal, Union, cast +from typing import Any, Literal, Union import matplotlib import matplotlib.patches as mpatches @@ -595,7 +595,7 @@ def _get_colors_for_categorical_obs( return palette[:len_cat] # type: ignore[return-value] -def _locate_points_value_in_table(value_key: str, sdata: SpatialData, element_name: str, table_name: str): +def _locate_points_value_in_table(value_key: str, sdata: SpatialData, table_name: str) -> _ValueOrigin: table = sdata[table_name] if value_key in table.obs.columns: @@ -608,7 +608,7 @@ def _locate_points_value_in_table(value_key: str, sdata: SpatialData, element_na # TODO consider move to relational query in spatialdata -def get_values_point_table(sdata: SpatialData, origin: _ValueOrigin, table_name: str): +def get_values_point_table(sdata: SpatialData, origin: _ValueOrigin, table_name: str) -> pd.Series: """Get a particular column stored in _ValueOrigin from the table in the spatialdata object.""" table = sdata[table_name] if origin.origin == "obs": @@ -624,8 +624,8 @@ def _set_color_source_vec( element_index: int, value_to_plot: str | None, element_name: list[str] | str | None = None, - groups: Sequence[str] | str | None = None, - palette: str | list[str] | None = None, + groups: Sequence[str | None] | str | None = None, + palette: list[str | None] | None = None, na_color: str | tuple[float, ...] | None = None, cmap_params: CmapParams | None = None, table_name: str | None = None, @@ -639,9 +639,7 @@ def _set_color_source_vec( # Figure out where to get the color from origins = _locate_value(value_key=value_to_plot, sdata=sdata, element_name=element_name, table_name=table_name) if model == PointsModel and table_name is not None: - origin = _locate_points_value_in_table( - value_key=value_to_plot, sdata=sdata, element_name=element_name, table_name=table_name - ) + origin = _locate_points_value_in_table(value_key=value_to_plot, sdata=sdata, table_name=table_name) if origin is not None: origins.append(origin) @@ -673,11 +671,17 @@ def _set_color_source_vec( color_source_vector = color_source_vector.remove_categories(categories.difference(groups)) categories = groups + palette_input: list[str] | str | None if groups is not None and groups[0] is not None: if isinstance(palette, list): - palette_input = palette[0] if palette[0] is None else palette + palette_input = ( + palette[0] + if palette[0] is None + else [color_palette for color_palette in palette if isinstance(color_palette, str)] + ) elif palette is not None and isinstance(palette, list): palette_input = palette[0] + else: palette_input = palette @@ -822,7 +826,7 @@ def _decorate_axs( ax: Axes, cax: PatchCollection, fig_params: FigParams, - value_to_plot: str | None, + value_to_plot: str | None, # str | None, color_source_vector: pd.Series[CategoricalDtype], adata: AnnData | None = None, palette: ListedColormap | str | list[str] | None = None, @@ -1351,11 +1355,13 @@ def _create_initial_element_table_mapping( ------- The updated render parameters. """ - element_table_mapping: dict[str, set[str]] = defaultdict(set) + element_table_mapping: dict[str, set[str | None] | str | None] = defaultdict(set) + if not params.element_table_mapping: for element_name in render_elements: - element_table_mapping[element_name].update(_get_element_annotators(sdata, element_name)) - else: + if isinstance(mapping := element_table_mapping[element_name], set): + mapping.update(_get_element_annotators(sdata, element_name)) + elif isinstance(params.element_table_mapping, (list, str)): table_names: list[str] = ( [params.element_table_mapping] if isinstance(params.element_table_mapping, str) @@ -1379,7 +1385,9 @@ def _create_initial_element_table_mapping( warnings.warn( f"The element '{element}' is not annotated by table '{table_name}'", UserWarning, stacklevel=2 ) - element_table_mapping[element].add(table_name) + if isinstance(mapping := element_table_mapping[element], set): + mapping.add(table_name) + assert isinstance(element_table_mapping, dict) params.element_table_mapping = element_table_mapping return params @@ -1387,38 +1395,46 @@ def _create_initial_element_table_mapping( def _update_element_table_mapping_label_colors( sdata: SpatialData, params: LabelsRenderParams | PointsRenderParams | ShapesRenderParams, render_elements: list[str] ) -> ImageRenderParams | LabelsRenderParams | PointsRenderParams | ShapesRenderParams: - element_table_mapping = params.element_table_mapping + element_table_mapping: dict[str, set[str | None] | str | None] | str | list[str] | None = ( + params.element_table_mapping + ) + + assert isinstance(element_table_mapping, dict) # If one color column check presence for each table annotating the specific element if isinstance(params.color, list) and len(params.color) == 1: params.color = params.color * len(render_elements) for element_name in render_elements: - for table_name in element_table_mapping[element_name].copy(): - if ( - params.color[0] not in sdata[table_name].obs.columns - and params.color[0] not in sdata[table_name].var_names - ): - element_table_mapping[element_name].remove(table_name) - if len(params.color) > 1: + if isinstance(mapping := element_table_mapping[element_name], set): + table_names = mapping.copy() + for table_name in table_names: + if ( + params.color[0] not in sdata[table_name].obs.columns + and params.color[0] not in sdata[table_name].var_names + ): + mapping.remove(table_name) + if isinstance(params.color, list) and len(params.color) > 1: assert len(params.color) == len( render_elements ), "Either one color should be given or the length should be equal to the number of elements being plotted." for index, element_name in enumerate(render_elements): - if len(element_table_mapping[element_name]) != 0: - for table_name in element_table_mapping[element_name].copy(): + if isinstance(mapping := element_table_mapping[element_name], set) and len(mapping) != 0: + for table_name in mapping.copy(): if ( params.color[index] not in sdata[table_name].obs.columns and params.color[index] not in sdata[table_name].var_names ): - element_table_mapping[element_name].remove(table_name) + mapping.remove(table_name) else: params.color[index] = None # We only want one table containing the color column per element + # table_set: set[str | None] for element_name, table_set in element_table_mapping.items(): - if len(table_set) > 1: + if isinstance(table_set, set) and len(table_set) > 1: raise ValueError(f"Multiple tables with color columns found for the element {element_name}") - element_table_mapping[element_name] = next(iter(table_set)) if len(table_set) != 0 else None + if isinstance(table_set, set): + element_table_mapping[element_name] = next(iter(table_set)) if len(table_set) != 0 else None params.element_table_mapping = element_table_mapping return params @@ -1427,8 +1443,13 @@ def _update_element_table_mapping_label_colors( def _validate_colors_element_table_mapping_points_shapes( sdata: SpatialData, params: PointsRenderParams | ShapesRenderParams, render_elements: list[str] ) -> PointsRenderParams | ShapesRenderParams: - element_table_mapping = cast(dict, params.element_table_mapping) - if len(params.color) == 1: + element_table_mapping: dict[str, set[str | None] | str | None] | str | list[str] | None = ( + params.element_table_mapping + ) + + assert isinstance(element_table_mapping, dict) + + if isinstance(params.color, list) and len(params.color) == 1 and isinstance(params.col_for_color, list): color = params.color[0] col_color = params.col_for_color[0] # This means that we are dealing with colors that are color like @@ -1444,13 +1465,13 @@ def _validate_colors_element_table_mapping_points_shapes( params.col_for_color.append(col_color) element_table_mapping[element_name] = set() else: - if len(element_table_mapping[element_name].copy()) != 0: - for table_name in element_table_mapping[element_name].copy(): + if isinstance(mapping := element_table_mapping[element_name], set) and len(mapping.copy()) != 0: + for table_name in mapping.copy(): if ( col_color not in sdata[table_name].obs.columns and col_color not in sdata[table_name].var_names ): - element_table_mapping[element_name].remove(table_name) + mapping.remove(table_name) params.col_for_color.append(None) else: params.col_for_color.append(col_color) @@ -1460,7 +1481,7 @@ def _validate_colors_element_table_mapping_points_shapes( params.color = [None] * len(render_elements) params.col_for_color = [None] * len(render_elements) else: - if len(params.color) != len(render_elements): + if isinstance(params.color, list) and len(params.color) != len(render_elements): warnings.warn( "The number of given colors and elements to render is not equal. " "Either provide one color or a list with one color for each element. skipping", @@ -1470,24 +1491,28 @@ def _validate_colors_element_table_mapping_points_shapes( params.color = [None] * len(render_elements) params.col_for_color = [None] * len(render_elements) else: + assert isinstance(params.color, list) + assert isinstance(params.col_for_color, list) for index, color in enumerate(params.color): if color is None: element_name = render_elements[index] col_color = params.col_for_color[index] - for table_name in element_table_mapping[element_name].copy(): - if ( - col_color not in sdata[table_name].obs.columns - and col_color not in sdata[table_name].var_names - and col_color not in sdata[element_name].columns - ): - element_table_mapping[element_name].remove(table_name) + if isinstance(mapping := element_table_mapping[element_name], set): + for table_name in mapping.copy(): + if ( + col_color not in sdata[table_name].obs.columns + and col_color not in sdata[table_name].var_names + and col_color not in sdata[element_name].columns + ): + mapping.remove(table_name) for index, element_name in enumerate(render_elements): # We only want one table value per element and only when there is a color column in the table if isinstance(params.col_for_color, list) and params.col_for_color[index] is not None: table_set = element_table_mapping[element_name] - if len(table_set) > 1: + if isinstance(table_set, set) and len(table_set) > 1: raise ValueError(f"More than one table found with color column {params.col_for_color[index]}.") - element_table_mapping[element_name] = next(iter(table_set)) if len(table_set) != 0 else None + if isinstance(tables := table_set, set): + element_table_mapping[element_name] = next(iter(tables)) if len(tables) != 0 else None if element_table_mapping[element_name] is None: warnings.warn( f"No table found with color column {params.col_for_color[index]} to render {element_name}", @@ -1606,7 +1631,7 @@ def _validate_render_params( alpha: float | int | None = None, channel: list[str] | list[int] | str | int | None = None, cmap: list[Colormap] | Colormap | str | None = None, - color: list[str] | str | None = None, + color: list[str | None] | str | None = None, contour_px: int | None = None, elements: list[str] | str | None = None, fill_alpha: float | int | None = None, @@ -1636,61 +1661,70 @@ def _validate_render_params( ) params_dict["elements"] = elements + groups_overwrite: list[list[str]] | None = None if groups is not None and element_type != "images": if not isinstance(groups, (list, str)): raise TypeError("Parameter 'groups' must be a string or a list of strings.") if isinstance(groups, str): - groups = [[groups]] + groups_overwrite = [[groups]] elif not isinstance(groups[0], list): - if not all(isinstance(g, str) for g in groups): + if all(isinstance(g, str) for g in groups): + groups_overwrite = [[group for group in groups if isinstance(group, str)]] + else: raise TypeError("All items in single 'groups' list must be strings.") - groups = [groups] + else: - if not all(isinstance(g, str) or g is None for group in groups for g in group): + if not all(isinstance(g, (str, type(None))) for group in groups for g in group): raise TypeError("All items in lists within lists of 'groups' must be strings or None.") - params_dict["groups"] = groups + params_dict["groups"] = groups_overwrite + palette_overwrite: list[list[str]] | None = None if palette is not None: if not isinstance(palette, (list, str)): raise TypeError("Parameter 'palette' must be a string or a list of strings.") if isinstance(palette, str): - palette = [[palette]] + palette_overwrite = [[palette]] elif not isinstance(palette[0], list): if not all(isinstance(pal, str) for pal in palette): raise TypeError("All items in single 'palette' list must be strings.") - palette = [palette] + palette_overwrite = [[pal for pal in palette if isinstance(pal, str)]] else: if not all(isinstance(p, str) or p is None for pal in palette for p in pal): raise TypeError("All items in lists within lists of 'groups' must be strings.") if element_type in ["shapes", "points", "labels"]: - if groups is None: + if groups_overwrite is None: raise ValueError("When specifying 'palette', 'groups' must also be specified.") - if len(groups) != len(palette): + if ( + groups_overwrite is not None + and palette_overwrite is not None + and len(groups_overwrite) != len(palette_overwrite) + ): raise ValueError( - f"The length of 'palette' and 'groups' must be the same, length is {len(palette)} and" - f"{len(groups)} respectively." + f"The length of 'palette' and 'groups' must be the same, length is {len(palette_overwrite)} and" + f"{len(groups_overwrite)} respectively." ) - for index, sublist in enumerate(groups): - if not len(sublist) == len(palette[index]): - raise ValueError("Not all nested lists in `groups` and `palette` are of equal length.") - if ( - not len(g_set := {type(el) for el in sublist}) - == len(p_set := {type(pal) for pal in palette[index]}) - == 1 - ): - raise ValueError( - "Mixed dtypes found in sublists of `groups` and/or `palette`. Must be either all" - "`str` or `None`." - ) - if g_set != p_set: - raise ValueError( - "Sublists with same index in `groups` and `palette` must contain elements of the " - "same dtype, either both `str` or `None`." - ) + if palette_overwrite is not None: + for index, sublist in enumerate(groups_overwrite): + if not len(sublist) == len(palette_overwrite[index]): + raise ValueError("Not all nested lists in `groups` and `palette` are of equal length.") + if ( + not len(g_set := {type(el) for el in sublist}) + == len(p_set := {type(pal) for pal in palette_overwrite[index]}) + == 1 + ): + raise ValueError( + "Mixed dtypes found in sublists of `groups` and/or `palette`. Must be either all" + "`str` or `None`." + ) + if g_set != p_set: + raise ValueError( + "Sublists with same index in `groups` and `palette` must contain elements of the " + "same dtype, either both `str` or `None`." + ) - params_dict["palette"] = palette + params_dict["palette"] = palette_overwrite if cmap is not None: if element_type == "images": @@ -1743,12 +1777,14 @@ def _validate_render_params( if not colors.is_color_like(outline_color): raise TypeError("Parameter 'outline_color' must be color-like.") + color_overwrite: list[str | None] = [] + col_for_color: list[str | None] if element_type in ["points", "shapes"]: - if color is not None: + if isinstance(color, (str, list)): if not isinstance(color, list): if colors.is_color_like(color): logger.info("Value for parameter 'color' appears to be a color, using it as such.") - color = [color] + color_overwrite = [color] col_for_color = [None] else: if not isinstance(color, str): @@ -1757,13 +1793,13 @@ def _validate_render_params( + "in sdata.table to use for coloring the shapes." ) col_for_color = [color] - color = [None] + color_overwrite = [None] else: col_for_color = [] - for index, c in enumerate(color): + for c in color: if colors.is_color_like(c): logger.info(f"Value `{c}` in list 'color' appears to be a color, using it as such.") - color[index] = c + color_overwrite.append(c) col_for_color.append(None) else: if not isinstance(c, str): @@ -1772,11 +1808,12 @@ def _validate_render_params( + "in sdata.table to use for coloring the shapes or should be color-like." ) col_for_color.append(c) - color[index] = None + color_overwrite.append(None) else: - color = [color] + color_overwrite = [color] col_for_color = [None] - params_dict["color"] = color + + params_dict["color"] = color_overwrite params_dict["col_for_color"] = col_for_color if element_type == "points": @@ -1839,23 +1876,26 @@ def _match_length_elements_groups_palette( params: ImageRenderParams | LabelsRenderParams | PointsRenderParams | ShapesRenderParams, render_elements: list[str], image: bool = False, -): +) -> ImageRenderParams | LabelsRenderParams | PointsRenderParams | ShapesRenderParams: if image and isinstance(params, ImageRenderParams): if params.palette is None: params.palette = [[None] for _ in range(len(render_elements))] else: params.palette = [params.palette[0] for _ in range(len(render_elements))] - else: + elif not isinstance(params, ImageRenderParams): groups = params.groups palette = params.palette + + groups_elements: list[list[str | None]] | None = None + palette_elements: list[list[str | None]] | None = None # We already checked before that length of groups and palette is the same if groups is not None: if len(groups) == 1: - params.groups = [groups[0] for _ in range(len(render_elements))] + groups_elements = [groups[0] for _ in range(len(render_elements)) if isinstance(groups[0], list)] if palette is not None: - params.palette = [palette[0] for _ in range(len(render_elements))] + palette_elements = [palette[0] for _ in range(len(render_elements)) if isinstance(palette[0], list)] else: - params.palette = [[None] for _ in range(len(render_elements))] + palette_elements = [[None] for _ in range(len(render_elements))] else: if len(groups) != len(render_elements): raise ValueError( @@ -1863,15 +1903,21 @@ def _match_length_elements_groups_palette( "of elements to be rendered." ) else: - params.groups = [[None] for _ in range(len(render_elements))] - params.palette = [[None] for _ in range(len(render_elements))] + groups_elements = [[None] for _ in range(len(render_elements))] + palette_elements = [[None] for _ in range(len(render_elements))] + params.palette = palette_elements + params.groups = groups_elements return params def _get_wanted_render_elements( - sdata, sdata_wanted_elements, params, cs, element_type: Literal["images", "labels", "points", "shapes"] -): + sdata: SpatialData, + sdata_wanted_elements: list[str], + params: ImageRenderParams | LabelsRenderParams | PointsRenderParams | ShapesRenderParams, + cs: str, + element_type: Literal["images", "labels", "points", "shapes"], +) -> tuple[list[str], list[str], bool]: wants_elements = True if element_type in ["images", "labels", "points", "shapes"]: # Prevents eval security risk wanted_elements = params.elements if params.elements is not None else list(getattr(sdata, element_type).keys()) @@ -1886,20 +1932,47 @@ def _get_wanted_render_elements( raise ValueError(f"Unknown element type {element_type}") -def _update_params(sdata, params, wanted_elements_on_cs, element_type: Literal["images", "labels", "points", "shapes"]): - if element_type in ["labels", "points", "shapes"] and wanted_elements_on_cs: +def _update_params( + sdata: SpatialData, + params: ImageRenderParams | LabelsRenderParams | PointsRenderParams | ShapesRenderParams, + wanted_elements_on_cs: list[str], + element_type: Literal["images", "labels", "points", "shapes"], +) -> ImageRenderParams | LabelsRenderParams | PointsRenderParams | ShapesRenderParams: + if isinstance(params, (LabelsRenderParams, PointsRenderParams, ShapesRenderParams)) and wanted_elements_on_cs: params = _create_initial_element_table_mapping(sdata, params, wanted_elements_on_cs) - if element_type == "labels": + if isinstance(params, LabelsRenderParams): params = _update_element_table_mapping_label_colors(sdata, params, wanted_elements_on_cs) - else: + if isinstance(params, (PointsRenderParams, ShapesRenderParams)): params = _validate_colors_element_table_mapping_points_shapes(sdata, params, wanted_elements_on_cs) - # if params.palette is None: - # params.palette = [[None] for _ in wanted_elements_on_cs] image_flag = element_type == "images" return _match_length_elements_groups_palette(params, wanted_elements_on_cs, image=image_flag) -def _is_coercable_to_float(series): +def _is_coercable_to_float(series: pd.Series) -> bool: numeric_series = pd.to_numeric(series, errors="coerce") return not numeric_series.isnull().any() + + +def _return_list_str_none(parameter: list[str | None] | str | None) -> list[str | None]: + """Force mypy to recognize list of string and None.""" + if isinstance(parameter, list) and all(isinstance(item, (str, type(None))) for item in parameter): + checked_parameter = parameter if isinstance(parameter, list) else [None] + else: + checked_parameter = [None] + return checked_parameter + + +def _return_list_list_str_none( + parameter: str | list[list[str | None]] | list[str | None] | None, +) -> list[list[str | None]]: + if not isinstance(parameter, list): + return [[None]] + + if all( + isinstance(sublist, list) and all(isinstance(inner_item, (str, type(None))) for inner_item in sublist) + for sublist in parameter + ): + return [list(sublist) for sublist in parameter if isinstance(sublist, list)] + + return [[None]] diff --git a/tests/conftest.py b/tests/conftest.py index 8ff5e953..ce1fb93e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -210,7 +210,6 @@ def sdata(request) -> SpatialData: s = SpatialData() else: s = request.getfixturevalue(request.param) - # print(f"request.param = {request.param}") return s