from __future__ import annotations
from typing import TYPE_CHECKING
import numpy as np
import pandas as pd
from mantispy._core._reduce import get_matrix
from mantispy._core.frames import as_frame
from mantispy._core.masks import reference_mask
from mantispy._core.plate import well_col, well_row
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
def _feature_values(adata: AnnData, feature: str | None) -> np.ndarray:
"""One number per row: a named feature, or the mean across features."""
matrix = get_matrix(adata)
if feature is not None:
return matrix[:, adata.var_names.get_loc(feature)].astype(float)
with np.errstate(invalid="ignore"):
return np.nanmean(matrix, axis=1)
[docs]
def plate_effects(adata: AnnData, feature: str | None = None, axes: np.ndarray | None = None) -> np.ndarray:
"""Row and column medians per plate, for spotting plate position artifacts.
Args:
adata: Object to draw, at any resolution.
feature: A single feature, or ``None`` for the mean across features.
axes: A ``(n_plates, 2)`` array of axes to draw into, or ``None`` for a new figure.
Returns:
The axes array, one row per plate, with the row marginal on the left and the column marginal on the right, each with the plate median drawn as a reference line.
Raises:
KeyError: ``feature`` is not one of ``var_names``, or ``obs`` has no ``Metadata_Plate`` or ``Metadata_Well`` column.
"""
import matplotlib.pyplot as plt
values = _feature_values(adata, feature)
frame = pd.DataFrame(
{
"plate": adata.obs["Metadata_Plate"].astype(str).to_numpy(),
"row": [well_row(well) for well in adata.obs["Metadata_Well"]],
"col": [well_col(well) for well in adata.obs["Metadata_Well"]],
"value": values,
}
)
plates = sorted(frame["plate"].unique())
if axes is None:
_, axes = plt.subplots(len(plates), 2, figsize=(9, 3 * len(plates)), squeeze=False)
for index, plate in enumerate(plates):
block = frame[frame["plate"] == plate]
reference = block["value"].median()
for position, axis_name in enumerate(("row", "col")):
axis = axes[index, position]
marginal = block.groupby(axis_name)["value"].median()
axis.plot(marginal.index, marginal.to_numpy(), marker="o", ms=3)
axis.axhline(reference, color="grey", ls="--", lw=1)
axis.set_xlabel(f"plate {axis_name}")
axis.set_title(f"{plate} by {axis_name}", fontsize=9)
axes[0, 0].set_ylabel(feature or "mean feature value")
return axes
[docs]
def image_qc(adata: AnnData, ax: Axes | None = None) -> Axes:
"""Image quality score per image, with the flagged images marked.
Args:
adata: Object :func:`~mantispy.pp.image_qc` has run on.
ax: Axes to draw on, or ``None`` for a new figure.
Returns:
The axes drawn on, with one point per image in the order of the table and the flagged images drawn larger and in crimson.
Raises:
KeyError: ``uns["mantispy"]`` holds no ``image_qc`` table.
"""
table = _table(adata, "image_qc", "mt.pp.image_qc")
ax = _axes(ax, (8, 4))
failed = ~table["qc_image_pass"].to_numpy(dtype=bool)
positions = np.arange(len(table))
ax.scatter(positions[~failed], table["qc_image_score"].to_numpy()[~failed], s=6, label="pass")
ax.scatter(positions[failed], table["qc_image_score"].to_numpy()[failed], s=18, color="crimson", label="flagged")
ax.set_xlabel("image")
ax.set_ylabel("quality score")
ax.legend(fontsize=7)
tidy = table.assign(image=positions, status=np.where(failed, "flagged", "pass"))
hover = [column for column in table.columns if column.startswith("Metadata_")]
_maybe_interactive(
"scatter",
ax=ax,
data=tidy,
x="image",
y="qc_image_score",
color="status",
hover=hover or None,
title="image quality",
)
return ax
[docs]
def control_drift(
adata: AnnData,
groupby: str = "Metadata_Plate",
n_components: int = 2,
ax: Axes | None = None,
) -> Axes:
"""Control wells projected onto principal components fitted on the controls alone.
Fitting on the controls alone shows how the reference moves between plates or batches, which is the drift normalization should remove.
Args:
adata: Object whose ``obs["Metadata_Control"]`` marks the wells to draw.
groupby: ``obs`` column that colors the control wells, normally the plate or the batch.
n_components: Components fitted on the controls.
The first two are the ones drawn.
ax: Axes to draw on, or ``None`` for a new figure.
Returns:
The axes drawn on, with one scatter per group of ``groupby`` in the space of the first two control components.
Raises:
KeyError: ``obs`` has no ``Metadata_Control`` column to select the controls with, or no ``groupby`` column.
ValueError: ``n_components`` is below the two that are drawn, or there are that many control rows or fewer, too few to fit them.
"""
from sklearn.decomposition import PCA
if n_components < 2:
raise ValueError(f"n_components must be at least 2, got {n_components}")
is_control = reference_mask(adata, "negcon")
if is_control.sum() < n_components + 1:
raise ValueError(f"need more than {n_components} control rows, found {int(is_control.sum())}")
controls = np.nan_to_num(get_matrix(adata)[is_control], nan=0.0, posinf=0.0, neginf=0.0)
embedding = PCA(n_components=n_components).fit_transform(controls)
labels = adata.obs[groupby].astype(str).to_numpy()[is_control]
ax = _axes(ax, (5, 4))
for group in pd.unique(labels):
selected = labels == group
ax.scatter(embedding[selected, 0], embedding[selected, 1], s=12, label=str(group))
ax.set_xlabel("control PC1")
ax.set_ylabel("control PC2")
ax.legend(title=groupby, fontsize=6, title_fontsize=7)
tidy = pd.DataFrame({"control PC1": embedding[:, 0], "control PC2": embedding[:, 1], groupby: labels})
_maybe_interactive(
"scatter", ax=ax, data=tidy, x="control PC1", y="control PC2", color=groupby, title="control drift"
)
return ax
[docs]
def outliers(
adata: AnnData, key: str = "qc_outlier", groupby: str = "Metadata_Plate", axes: np.ndarray | None = None
) -> np.ndarray:
"""Outlier score distribution, and the flagged fraction per ``groupby`` group.
Args:
adata: Object :func:`~mantispy.pp.outliers` has run on.
key: ``obs`` column holding the flag, whose score is read from ``key + "_score"``.
groupby: ``obs`` column whose groups become the bars, e.g. ``"Metadata_Well"`` on a single plate.
axes: A pair of axes to draw into, or ``None`` for a new figure.
Returns:
The two axes: the score histogram split into kept and flagged, and the flagged fraction per group.
Raises:
KeyError: ``obs`` has no ``key`` column.
"""
import matplotlib.pyplot as plt
if key not in adata.obs:
raise KeyError(f"obs has no {key!r}; run mt.pp.outliers first")
if axes is None:
_, axes = plt.subplots(1, 2, figsize=(9, 3.5), layout="constrained")
scores = as_frame(adata.obs)[f"{key}_score"].to_numpy(dtype=float)
flagged = as_frame(adata.obs)[key].to_numpy(dtype=bool)
axes[0].hist(scores[~flagged], bins=50, label="kept")
axes[0].hist(scores[flagged], bins=50, color="crimson", label="flagged")
axes[0].set_xlabel("outlier score")
axes[0].legend(fontsize=7)
per_group = as_frame(adata.obs).groupby(groupby, observed=True)[key].mean()
axes[1].bar(np.arange(len(per_group)), per_group.to_numpy())
axes[1].set_xticks(np.arange(len(per_group)))
axes[1].set_xticklabels([str(name) for name in per_group.index], rotation=45, fontsize=7)
axes[1].set_xlabel(groupby)
axes[1].set_ylabel("fraction flagged")
return axes