"""Partition and metacell visualization for AnnData."""
from __future__ import annotations
import matplotlib.pyplot as plt
import pandas as pd
import seaborn as sns
from anndata import AnnData
from matplotlib.figure import Figure
from campmc.mt.sizes import metacell_sizes
from campmc.pl._helpers import _auto_palette, _embedding_2d, _finalize_figure
[docs]
def assignment(
adata: AnnData,
*,
embedding_key: str = "X_umap",
partition_key: str = "camp",
colour_metacells: bool = True,
title: str = "Metacell Assignments",
figsize: tuple[float, float] = (5.0, 5.0),
metacell_size: float = 20.0,
cell_size: float = 10.0,
show: bool = False,
) -> Figure | None:
"""Plot a 2D embedding coloured by CAMP metacell assignments.
Parameters
----------
adata
AnnData with partition labels in ``obs`` and a 2D embedding in ``obsm``.
embedding_key
``obsm`` key for the 2D plot (e.g. ``X_umap`` or ``X_pca``).
partition_key
Column in ``adata.obs`` with metacell labels.
colour_metacells
If ``True``, colour cells by metacell and overlay group-mean centroids.
If ``False``, grey cells with red centroid markers.
title
Figure title.
figsize
Figure size in inches.
metacell_size
Marker size for metacell centroids.
cell_size
Marker size for single cells.
show
If ``True``, display via ``plt.show()`` and return ``None``.
Returns
-------
The figure when ``show`` is false; otherwise ``None`` after displaying.
Examples
--------
.. exec-jupyter::
import campmc as cp
import scanpy as sc
adata = sc.datasets.pbmc68k_reduced()
cp.partition(adata, method="camp3", gamma=50, random_state=0)
cp.pl.assignment(adata, embedding_key="X_umap", show=True)
"""
if partition_key not in adata.obs.columns:
raise ValueError(
f"adata.obs is missing partition column '{partition_key}'. "
"Run campmc.partition before plotting assignments."
)
coords = _embedding_2d(adata, embedding_key)
plot_df = coords.join(adata.obs[[partition_key]])
plot_df[partition_key] = plot_df[partition_key].astype("category")
centroids = plot_df.groupby(partition_key, observed=True).mean(numeric_only=True).reset_index()
fig, ax = plt.subplots(figsize=figsize)
palette = _auto_palette(plot_df[partition_key].nunique())
x_col, y_col = coords.columns[0], coords.columns[1]
if colour_metacells:
sns.scatterplot(
data=plot_df,
x=x_col,
y=y_col,
hue=partition_key,
palette=palette,
s=cell_size,
linewidth=0,
legend=False,
ax=ax,
)
sns.scatterplot(
data=centroids,
x=x_col,
y=y_col,
hue=partition_key,
palette=palette,
s=metacell_size,
edgecolor="black",
linewidth=1.25,
legend=False,
ax=ax,
)
else:
sns.scatterplot(
data=plot_df,
x=x_col,
y=y_col,
color="grey",
s=cell_size,
legend=False,
ax=ax,
)
sns.scatterplot(
data=centroids,
x=x_col,
y=y_col,
color="red",
s=metacell_size,
edgecolor="black",
linewidth=1.25,
legend=False,
ax=ax,
)
ax.set_xlabel(x_col)
ax.set_ylabel(y_col)
ax.set_title(title)
ax.set_axis_off()
fig.tight_layout()
return _finalize_figure(fig, show=show)
[docs]
def sizes(
adata: AnnData,
*,
partition_key: str = "camp",
bins: int | None = None,
title: str = "Distribution of Metacell Sizes",
figsize: tuple[float, float] = (5.0, 5.0),
show: bool = False,
) -> tuple[Figure | None, pd.DataFrame]:
"""Plot the distribution of cells per metacell.
Parameters
----------
adata
AnnData with partition labels in ``obs``.
partition_key
Column in ``adata.obs`` with metacell labels.
bins
Number of histogram bins; ``None`` uses seaborn defaults.
title
Figure title.
figsize
Figure size in inches.
show
If ``True``, display via ``plt.show()`` and return ``None`` for the figure.
Returns
-------
Tuple of ``(figure, sizes)`` where *sizes* is from {func}`campmc.mt.metacell_sizes`.
*figure* is ``None`` when ``show`` displays the plot.
Examples
--------
.. exec-jupyter::
import campmc as cp
import scanpy as sc
adata = sc.datasets.pbmc68k_reduced()
cp.partition(adata, method="camp3", gamma=50, random_state=0)
cp.pl.sizes(adata, show=True)
"""
size_df = metacell_sizes(adata, partition_key=partition_key)
fig, ax = plt.subplots(figsize=figsize)
hist_kwargs: dict = {"kde": False, "ax": ax}
if bins is not None:
hist_kwargs["bins"] = bins
sns.histplot(size_df["size"], **hist_kwargs)
sns.despine(ax=ax)
ax.set_xlabel("Number of Cells per Metacell")
ax.set_title(title)
fig.tight_layout()
return _finalize_figure(fig, show=show), size_df