Source code for mantispy.pl._qc

"""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 returned as _returned
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 | None: """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 when the caller passed ``ax``, else ``None`` because the plot then owns the figure it created. 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 _returned(ax)
[docs] def feature_distributions( adata: AnnData, features: Sequence[str], groupby: str = "Metadata_Plate", layer_before: str | None = "raw", kind: str = "ecdf", ) -> np.ndarray | None: """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: ``None``; this plot always creates its own figure, a 2-D grid 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 _returned(axes, owned=True)
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 | None: """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 when the caller passed ``ax``, else ``None`` because the plot then owns the figure it created. """ 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 _returned(ax)
[docs] def qc(adata: AnnData, figsize: tuple[float, float] = (12, 8)) -> np.ndarray | None: """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: ``None``; this plot always creates its own figure, a two-by-two grid of the QC panels. """ 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 _returned(axes, owned=True)
[docs] def replicate_saturation(adata: AnnData, key: str = "replicate_saturation", ax: Axes | None = None) -> Axes | None: """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 when the caller passed ``ax``, else ``None`` because the plot then owns the figure it created. 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 _returned(ax)
[docs] def cytotoxicity(adata: AnnData, key: str = "cytotoxicity", label_top: int = 8, ax: Axes | None = None) -> Axes | None: """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 when the caller passed ``ax``, else ``None`` because the plot then owns the figure it created. 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 _returned(ax)