"""Quality-control plots."""
from __future__ import annotations
from collections.abc import Sequence
from typing import TYPE_CHECKING
import numpy as np
import pandas as pd
from mantispy._core._reduce import get_matrix, group_codes
from mantispy._core.frames import as_frame
from mantispy._core.schema import get_resolution
from mantispy.pl._common import axes as _axes
from mantispy.pl._common import maybe_interactive as _maybe_interactive
from mantispy.pl._common import table as _table
if TYPE_CHECKING:
from anndata import AnnData
from matplotlib.axes import Axes
[docs]
def cell_counts(
adata: AnnData, groupby: str = "Metadata_Plate", ax: Axes | None = None, count_key: str = "Metadata_CellCount"
) -> Axes:
"""Distribution of cells per well, split by ``groupby``.
Args:
adata: Cells, which are counted per well, or profiles, whose ``count_key`` is drawn.
groupby: ``obs`` column whose groups become the boxes.
ax: Axes to draw on, or ``None`` for a new figure.
count_key: ``obs`` column holding the cell count of profiles.
Returns:
The axes drawn on.
Raises:
KeyError: Profiles carry no ``count_key``.
"""
obs = as_frame(adata.obs)
cells = get_resolution(adata) == "cell"
if cells:
codes, keys = group_codes(adata, ["Metadata_Plate", "Metadata_Well"])
counts = np.bincount(codes, minlength=len(keys))
labels = obs.groupby(codes, observed=True)[groupby].first()
elif count_key in obs:
counts = obs[count_key].to_numpy(dtype=float)
labels = obs[groupby]
else:
raise KeyError(f"obs has no column {count_key!r}; mt.tl.aggregate writes one, or name another with count_key=")
ax = _axes(ax, (6, 4))
# A missing label is a group of its own rather than one that matches nothing.
groups, names = labels.factorize(use_na_sentinel=False)
known = np.isfinite(counts)
ax.boxplot([counts[known & (groups == i)] for i in range(len(names))], tick_labels=[str(name) for name in names])
ylabel = "cells per well" if cells else count_key
ax.set_ylabel(ylabel)
ax.set_xlabel(groupby)
ax.tick_params(axis="x", rotation=45)
label_of = np.array([str(name) for name in names])
tidy = pd.DataFrame({groupby: label_of[groups[known]], ylabel: counts[known]})
_maybe_interactive("box", ax=ax, data=tidy, x=groupby, y=ylabel, title="cell counts")
return ax
[docs]
def feature_distributions(
adata: AnnData,
features: Sequence[str],
groupby: str = "Metadata_Plate",
layer_before: str | None = "raw",
kind: str = "ecdf",
) -> np.ndarray:
"""Per-feature distributions, before and after normalization when ``layer_before`` exists.
Args:
adata: Object holding the features to draw.
features: ``var_names`` to draw, one column of panels each.
groupby: ``obs`` column whose groups are drawn separately within each panel.
layer_before: Layer holding the values before normalization, or ``None`` to draw only the current ones.
A layer that the object does not hold is skipped in the same way.
kind: ``"ecdf"``, ``"hist"``, or ``"ridge"`` for one offset filled density per group, which is easier to read with many groups.
Returns:
A 2-D array of axes with one row per layer shown and one column per feature.
Raises:
ValueError: ``kind`` is not one of the three accepted values.
"""
import matplotlib.pyplot as plt
if kind not in {"ecdf", "hist", "ridge"}:
raise ValueError(f"kind must be 'ecdf', 'hist' or 'ridge', got {kind!r}")
features = list(features)
show_before = layer_before is not None and layer_before in adata.layers
layers = [layer_before, None] if show_before else [None]
figure, axes = plt.subplots(len(layers), len(features), figsize=(4 * len(features), 3 * len(layers)), squeeze=False)
for row, layer in enumerate(layers):
matrix = get_matrix(adata, layer)
for column, feature in enumerate(features):
axis = axes[row, column]
values = matrix[:, adata.var_names.get_loc(feature)]
for offset, group in enumerate(dict.fromkeys(adata.obs[groupby])):
selected = values[(adata.obs[groupby] == group).to_numpy()]
selected = selected[~np.isnan(selected)]
if not selected.size:
continue
if kind == "ecdf":
axis.plot(np.sort(selected), np.linspace(0, 1, selected.size), lw=1, label=str(group))
elif kind == "ridge":
_ridge(axis, selected, offset, str(group))
else:
axis.hist(selected, bins=50, histtype="step", density=True, label=str(group))
axis.set_title(f"{feature}\n{'raw' if layer else 'current'}", fontsize=8)
axes[0, 0].legend(fontsize=6)
figure.tight_layout()
return axes
def _ridge(axis: Axes, values: np.ndarray, offset: int, label: str) -> None:
"""One filled density curve, raised by ``offset`` so the groups stack rather than overlap."""
grid = np.linspace(values.min(), values.max(), 128)
if values.size < 2 or np.ptp(values) == 0:
return
from scipy.stats import gaussian_kde
density = gaussian_kde(values)(grid)
density = density / density.max() * 0.9
axis.fill_between(grid, offset, offset + density, alpha=0.7, lw=0.6, edgecolor="black", label=label)
[docs]
def nan_matrix(adata: AnnData, max_features: int = 200, ax: Axes | None = None) -> Axes:
"""Fraction of missing values per feature, per plate.
Args:
adata: Object to measure the missing values of.
max_features: How many features to draw, taken in ``var_names`` order.
ax: Axes to draw on, or ``None`` for a new figure.
Returns:
The axes drawn on.
"""
ax = _axes(ax, (8, 4))
missing = np.isnan(get_matrix(adata))
codes, keys = group_codes(adata, "Metadata_Plate")
fractions = np.stack([missing[codes == index].mean(axis=0) for index in range(len(keys))])
shown = fractions[:, :max_features]
image = ax.imshow(shown, aspect="auto", cmap="magma", vmin=0, vmax=1)
ax.set_yticks(range(len(keys)))
ax.set_yticklabels([str(key) for key in keys], fontsize=7)
ax.set_xlabel("feature")
ax.figure.colorbar(image, ax=ax, label="NaN fraction")
_maybe_interactive(
"heatmap",
ax=ax,
matrix=shown,
rows=[str(key) for key in keys],
columns=[str(name) for name in adata.var_names[:max_features]],
value_label="NaN fraction",
title="missing values",
)
return ax
[docs]
def qc(adata: AnnData, figsize: tuple[float, float] = (12, 8)) -> np.ndarray:
"""Two-by-two summary of the QC metrics :func:`~mantispy.pp.calculate_qc_metrics` writes.
Each panel is drawn only if the object holds what it needs, so a partial run still gives a figure.
Args:
adata: Object carrying the QC annotations.
figsize: Size of the whole figure, in inches.
Returns:
The two-by-two array of axes.
"""
import matplotlib.pyplot as plt
figure, axes = plt.subplots(2, 2, figsize=figsize)
if get_resolution(adata) == "cell" or "Metadata_CellCount" in adata.obs:
cell_counts(adata, ax=axes[0, 0])
if "qc_n_nan" in adata.var:
axes[0, 1].hist(adata.var["qc_n_nan"], bins=40)
axes[0, 1].set_xlabel("cells missing this feature")
flags = [column for column in ("qc_is_border", "qc_area_outlier", "qc_pass") if column in adata.obs]
if flags:
obs = as_frame(adata.obs)
obs.groupby("Metadata_Plate", observed=True)[flags].mean().plot.bar(ax=axes[1, 0])
axes[1, 0].set_ylabel("fraction of cells")
axes[1, 0].legend(fontsize=6)
if "qc_variance" in adata.var:
variance = as_frame(adata.var)["qc_variance"].to_numpy(dtype=float)
variance = variance[np.isfinite(variance) & (variance > 0)]
if variance.size:
axes[1, 1].hist(np.log10(variance), bins=40)
axes[1, 1].set_xlabel("log10 feature variance")
figure.tight_layout()
return axes
[docs]
def replicate_saturation(adata: AnnData, key: str = "replicate_saturation", ax: Axes | None = None) -> Axes:
"""The saturation curve with its spread across draws.
A curve still rising at the right edge means the screen is under-replicated, which informs the design of the next experiment.
Args:
adata: Object holding the table :func:`~mantispy.tl.replicate_saturation` wrote.
key: Name of that table in ``uns["mantispy"]``.
ax: Axes to draw on, or ``None`` for a new figure.
Returns:
The axes drawn on.
Raises:
KeyError: ``uns["mantispy"]`` holds no table under ``key``.
"""
table = _table(adata, key, "mt.tl.replicate_saturation")
ax = _axes(ax, (5, 4))
ax.errorbar(table["n_replicates"], table["mean"], yerr=table["std"], marker="o", capsize=3)
ax.set_xticks(table["n_replicates"].to_numpy())
ax.set_xlabel("replicates per perturbation")
ax.set_ylabel("signature agreement")
tidy = pd.DataFrame(
{
"replicates per perturbation": table["n_replicates"].to_numpy(dtype=float),
"signature agreement": table["mean"].to_numpy(dtype=float),
"std": table["std"].to_numpy(dtype=float),
}
)
_maybe_interactive(
"line",
ax=ax,
data=tidy,
x="replicates per perturbation",
y="signature agreement",
hover=["std"],
title="replicate saturation",
)
return ax
[docs]
def cytotoxicity(adata: AnnData, key: str = "cytotoxicity", label_top: int = 8, ax: Axes | None = None) -> Axes:
"""Distance from the controls against viability, with the suspect groups marked.
Groups in the upper left are far from the controls and have lost most of their cells.
The dashed line is the minimum viability the run used.
Args:
adata: Object holding the table :func:`~mantispy.tl.cytotoxicity` wrote.
key: Name of that table in ``uns["mantispy"]``.
label_top: How many of the most distant suspect groups to label.
ax: Axes to draw on, or ``None`` for a new figure.
Returns:
The axes drawn on.
Raises:
KeyError: ``uns["mantispy"]`` holds no table under ``key``.
"""
table = _table(adata, key, "mt.tl.cytotoxicity")
suspect = table["suspect"].to_numpy(dtype=bool)
ax = _axes(ax, (5.5, 4.5))
ax.scatter(table["viability"][~suspect], table["distance"][~suspect], s=16, color="tab:blue", label="ok")
ax.scatter(table["viability"][suspect], table["distance"][suspect], s=16, color="crimson", label="suspect")
threshold = adata.uns.get("mantispy", {}).get("params", {}).get("cytotoxicity", {}).get("min_viability", 0.7)
ax.axvline(threshold, color="grey", ls="--", lw=1, label=f"viability = {threshold}")
for _, row in table[suspect].nlargest(label_top, "distance").iterrows():
ax.annotate(str(row["group"]), (row["viability"], row["distance"]), fontsize=6)
ax.set_xlabel("viability, relative to the controls")
ax.set_ylabel("distance from the controls")
ax.legend(fontsize=7)
tidy = pd.DataFrame(
{
"group": table["group"].astype(str).to_numpy(),
"viability, relative to the controls": table["viability"].to_numpy(dtype=float),
"distance from the controls": table["distance"].to_numpy(dtype=float),
"status": np.where(suspect, "suspect", "ok"),
}
)
_maybe_interactive(
"scatter",
ax=ax,
data=tidy,
x="viability, relative to the controls",
y="distance from the controls",
color="status",
hover=["group"],
title="cytotoxicity",
)
return ax