diff --git a/CHANGELOG.md b/CHANGELOG.md index ea577174..955a092c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -19,6 +19,7 @@ and this project adheres to [Semantic Versioning][]. ### Fixed - Now dropping index when plotting shapes after spatial query (#177) +- User can now pass Colormap objects to the cmap argument in render_images. When only one cmap is given for 3 channels, it is now applied to each channel (#188, #194) ## [0.0.6] - 2023-11-06 diff --git a/src/spatialdata_plot/pl/basic.py b/src/spatialdata_plot/pl/basic.py index c4352f05..f5b8375b 100644 --- a/src/spatialdata_plot/pl/basic.py +++ b/src/spatialdata_plot/pl/basic.py @@ -362,9 +362,6 @@ def render_images( sdata = _verify_plotting_tree(sdata) n_steps = len(sdata.plotting_tree.keys()) - if channel is None and cmap is None: - cmap = "brg" - cmap_params: list[CmapParams] | CmapParams if isinstance(cmap, list): cmap_params = [ diff --git a/src/spatialdata_plot/pl/render.py b/src/spatialdata_plot/pl/render.py index 03d95662..ffb83eab 100644 --- a/src/spatialdata_plot/pl/render.py +++ b/src/spatialdata_plot/pl/render.py @@ -434,10 +434,29 @@ def _render_images( if render_params.cmap_params[i].norm is not None: layers[c] = render_params.cmap_params[i].norm(layers[c]) - # 2A) Image has 3 channels, no palette/cmap info -> use RGB - if n_channels == 3 and render_params.palette is None and not got_multiple_cmaps: + # 2A) Image has 3 channels, no palette info, and no/only one cmap was given + if n_channels == 3 and render_params.palette 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 + channel_cmaps = [render_params.cmap_params.cmap] * n_channels + # Apply cmaps to each channel, add up and normalize to [0, 1] + stacked = ( + np.stack([channel_cmaps[i](layers[c]) for i, c in enumerate(channels)], 0).sum(0) / n_channels + ) + # Remove alpha channel so we can overwrite it from render_params.alpha + stacked = stacked[:, :, :3] + logger.warning( + "One cmap was given for multiple channels and is now used for each channel. " + "You're blending multiple cmaps. " + "If the plot doesn't look like you expect, it might be because your " + "cmaps go from a given color to 'white', and not to 'transparent'. " + "Therefore, the 'white' of higher layers will overlay the lower layers. " + "Consider using 'palette' instead." + ) + im = ax.imshow( - np.stack([layers[c] for c in channels], axis=-1), + stacked, alpha=render_params.alpha, ) im.set_transform(trans_data) diff --git a/src/spatialdata_plot/pl/utils.py b/src/spatialdata_plot/pl/utils.py index f85e0144..7eabe400 100644 --- a/src/spatialdata_plot/pl/utils.py +++ b/src/spatialdata_plot/pl/utils.py @@ -344,7 +344,13 @@ def _prepare_cmap_norm( **kwargs: Any, ) -> CmapParams: is_default = cmap is None - cmap = copy(matplotlib.colormaps[rcParams["image.cmap"] if cmap is None else cmap]) + if cmap is None: + cmap = rcParams["image.cmap"] + if isinstance(cmap, str): + cmap = matplotlib.colormaps[cmap] + + cmap = copy(cmap) + cmap.set_bad("lightgray" if na_color is None else na_color) if isinstance(norm, Normalize) or not norm: diff --git a/tests/_images/Images_can_pass_cmap.png b/tests/_images/Images_can_pass_cmap.png new file mode 100644 index 00000000..c8309350 Binary files /dev/null and b/tests/_images/Images_can_pass_cmap.png differ diff --git a/tests/_images/Images_can_pass_cmap_list.png b/tests/_images/Images_can_pass_cmap_list.png new file mode 100644 index 00000000..1a37611c Binary files /dev/null and b/tests/_images/Images_can_pass_cmap_list.png differ diff --git a/tests/_images/Images_can_pass_cmap_to_render_images.png b/tests/_images/Images_can_pass_cmap_to_render_images.png deleted file mode 100644 index f18fa5f0..00000000 Binary files a/tests/_images/Images_can_pass_cmap_to_render_images.png and /dev/null differ diff --git a/tests/_images/Images_can_pass_str_cmap.png b/tests/_images/Images_can_pass_str_cmap.png new file mode 100644 index 00000000..c8309350 Binary files /dev/null and b/tests/_images/Images_can_pass_str_cmap.png differ diff --git a/tests/_images/Images_can_pass_str_cmap_list.png b/tests/_images/Images_can_pass_str_cmap_list.png new file mode 100644 index 00000000..1a37611c Binary files /dev/null and b/tests/_images/Images_can_pass_str_cmap_list.png differ diff --git a/tests/pl/test_render_images.py b/tests/pl/test_render_images.py index a1b732ff..57667c8e 100644 --- a/tests/pl/test_render_images.py +++ b/tests/pl/test_render_images.py @@ -24,9 +24,21 @@ class TestImages(PlotTester, metaclass=PlotTesterMeta): def test_plot_can_render_image(self, sdata_blobs: SpatialData): sdata_blobs.pl.render_images(elements="blobs_image").pl.show() - def test_plot_can_pass_cmap_to_render_images(self, sdata_blobs: SpatialData): + def test_plot_can_pass_str_cmap(self, sdata_blobs: SpatialData): sdata_blobs.pl.render_images(elements="blobs_image", cmap="seismic").pl.show() + def test_plot_can_pass_cmap(self, sdata_blobs: SpatialData): + sdata_blobs.pl.render_images(elements="blobs_image", cmap=matplotlib.colormaps["seismic"]).pl.show() + + def test_plot_can_pass_str_cmap_list(self, sdata_blobs: SpatialData): + sdata_blobs.pl.render_images(elements="blobs_image", cmap=["seismic", "Reds", "Blues"]).pl.show() + + def test_plot_can_pass_cmap_list(self, sdata_blobs: SpatialData): + sdata_blobs.pl.render_images( + elements="blobs_image", + cmap=[matplotlib.colormaps["seismic"], matplotlib.colormaps["Reds"], matplotlib.colormaps["Blues"]], + ).pl.show() + def test_plot_can_render_a_single_channel_from_image(self, sdata_blobs: SpatialData): sdata_blobs.pl.render_images(elements="blobs_image", channel=0).pl.show()