"""Plots for hit calling, effect sizes and dose response."""
from __future__ import annotations
from typing import TYPE_CHECKING, Any
import numpy as np
import pandas as pd
from mantispy._core.frames import as_frame
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
from mantispy.tl._dose import DOSE_PHASES
if TYPE_CHECKING:
from anndata import AnnData
from matplotlib.axes import Axes
_FLOOR = 1e-12
def _significance(values: pd.Series | np.ndarray) -> np.ndarray:
"""Q-values on a ``-log10`` scale, floored so that an exact zero stays on the plot."""
return -np.log10(np.clip(np.asarray(values, dtype=float), _FLOOR, None))
def _threshold(adata: AnnData, function: str, default: float = 0.05) -> float:
"""The threshold the run used, so the drawn line matches the table."""
recorded = adata.uns.get("mantispy", {}).get("params", {}).get(function, {})
return float(recorded.get("threshold", default))
def _count_either_side(ax: Axes, significance: np.ndarray, line: float) -> None:
"""Write how many points sit above and below a horizontal line, next to it at the right edge.
Above means strictly above, as a hit is strictly below the q-value threshold.
"""
shown = significance[np.isfinite(significance)]
above = int((shown > line).sum())
for text, offset, align in ((f"{above} above", 2, "bottom"), (f"{shown.size - above} below", -2, "top")):
ax.annotate(
text,
(0.98, line),
xycoords=ax.get_yaxis_transform(),
xytext=(0, offset),
textcoords="offset points",
ha="right",
va=align,
fontsize=7,
color="0.35",
)
def _recorded_response(adata: AnnData, function: str, default: str) -> str:
"""The response column the run read, so the plot draws what the table was built from."""
recorded = adata.uns.get("mantispy", {}).get("params", {}).get(function, {})
return str(recorded.get("response", default))
[docs]
def hits(adata: AnnData, key: str = "hits", label_top: int = 10, ax: Axes | None = None) -> Axes:
"""Distance from the controls against significance, with the most distant groups labeled.
A point in the upper right moved far from the controls and is significant under the permutation null.
The dashed line is the q-value threshold the run used, so every point colored as a hit sits on or above it, and the counts beside it say how many groups sit on either side.
Args:
adata: Object holding the table :func:`~mantispy.tl.hit_calling` wrote.
key: Name of that table in ``uns["mantispy"]``.
label_top: How many of the most distant 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.hit_calling")
ax = _axes(ax, (5.5, 4.5))
significance = _significance(table["qvalue"])
called = table["is_hit"].to_numpy(dtype=bool)
ax.scatter(table["distance"][~called], significance[~called], s=14, color="lightgrey", label="not called")
ax.scatter(table["distance"][called], significance[called], s=14, color="crimson", label="hit")
# Params are stored under the function name, not the table's `key`.
threshold = _threshold(adata, "hit_calling")
ax.axhline(-np.log10(threshold), color="grey", ls="--", lw=1, label=f"q = {threshold}")
_count_either_side(ax, significance, -np.log10(threshold))
for _, row in table.nlargest(label_top, "distance").iterrows():
ax.annotate(str(row["group"]), (row["distance"], -np.log10(max(float(row["qvalue"]), _FLOOR))), fontsize=6)
ax.set_xlabel("distance from the controls")
ax.set_ylabel("-log10 q")
ax.legend(fontsize=7)
tidy = pd.DataFrame(
{
"group": table["group"].astype(str).to_numpy(),
"distance from the controls": table["distance"].to_numpy(dtype=float),
"-log10 q": significance,
"called": np.where(called, "hit", "not called"),
}
)
_maybe_interactive(
"scatter",
ax=ax,
data=tidy,
x="distance from the controls",
y="-log10 q",
color="called",
hover=["group"],
title="hits",
)
return ax
def _rows_for(table: pd.DataFrame, column: str, value: str, key: str) -> pd.DataFrame:
"""The table's rows for one group, or a KeyError naming the groups it does hold."""
selected = table[table[column].astype(str) == str(value)]
if selected.empty:
raise KeyError(f"no {column} {value!r} in uns['mantispy'][{key!r}]; it holds {sorted(set(table[column]))[:5]}")
return selected
def _effects(adata: AnnData, group: str, key: str) -> tuple[pd.DataFrame, pd.Series | None]:
table = _table(adata, key, "mt.tl.effect_size")
selected = _rows_for(table, "group", group, key)
var = as_frame(adata.var)
families = var["feature_group"].astype(str) if "feature_group" in var else None
return selected, families
def _family_colours(families: pd.Series | None, names: pd.Series) -> tuple[list[Any], dict[str, Any]]:
"""One color per feature family, assigned in sorted order of the families present."""
import matplotlib.pyplot as plt
if families is None:
return ["tab:blue"] * len(names), {}
labels = [str(families.get(name, "unknown")) for name in names]
palette = {label: plt.get_cmap("tab20")(index % 20) for index, label in enumerate(sorted(set(labels)))}
return [palette[label] for label in labels], palette
[docs]
def effect_sizes(adata: AnnData, group: str, key: str = "effect", top: int = 30, ax: Axes | None = None) -> Axes:
"""The largest effects for one group, colored by feature family.
Args:
adata: Object holding the table :func:`~mantispy.tl.effect_size` wrote.
group: Which group of that table to draw.
key: Name of that table in ``uns["mantispy"]``.
top: How many features to draw, taken by absolute effect.
ax: Axes to draw on, or ``None`` for a new figure.
Returns:
The axes drawn on.
Raises:
KeyError: There is no such table, or it holds no such group.
"""
import matplotlib.pyplot as plt
selected, families = _effects(adata, group, key)
strongest = selected.reindex(selected["effect"].abs().sort_values(ascending=False).index).head(top)[::-1]
ax = _axes(ax, (6, 0.22 * len(strongest) + 1.5))
colours, palette = _family_colours(families, strongest["feature"])
ax.barh(np.arange(len(strongest)), strongest["effect"], color=colours)
ax.set_yticks(np.arange(len(strongest)))
ax.set_yticklabels(strongest["feature"], fontsize=6)
ax.axvline(0, color="black", lw=0.8)
ax.set_xlabel("effect size")
ax.set_title(str(group), fontsize=9)
if palette:
ax.legend(
handles=[plt.Line2D([], [], color=colour, lw=6, label=label) for label, colour in palette.items()],
fontsize=6,
loc="lower right",
)
family_of = None if families is None else [str(families.get(name, "unknown")) for name in strongest["feature"]]
tidy = pd.DataFrame(
{
"feature": strongest["feature"].astype(str).to_numpy(),
"effect size": strongest["effect"].to_numpy(dtype=float),
}
)
if family_of is not None:
tidy["family"] = family_of
_maybe_interactive(
"barh",
ax=ax,
data=tidy,
x="effect size",
y="feature",
color="family" if family_of is not None else None,
title=str(group),
)
return ax
[docs]
def feature_volcano(
adata: AnnData, group: str, key: str = "effect", label_top: int = 8, ax: Axes | None = None
) -> Axes:
"""Effect against significance, per feature, for one group.
The dashed line is q = 0.05, and the counts beside it say how many features sit on either side.
Args:
adata: Object holding the table :func:`~mantispy.tl.effect_size` wrote.
group: Which group of that table to draw.
key: Name of that table in ``uns["mantispy"]``.
label_top: How many features to label, taken by absolute effect.
ax: Axes to draw on, or ``None`` for a new figure.
Returns:
The axes drawn on.
Raises:
KeyError: There is no such table, or it holds no such group.
"""
import matplotlib.pyplot as plt
selected, families = _effects(adata, group, key)
ax = _axes(ax, (5.5, 4.5))
colours, palette = _family_colours(families, selected["feature"])
significance = _significance(selected["qvalue"])
ax.scatter(selected["effect"], significance, s=12, color=colours)
for _, row in (
selected.reindex(selected["effect"].abs().sort_values(ascending=False).index).head(label_top).iterrows()
):
ax.annotate(str(row["feature"]), (row["effect"], -np.log10(max(float(row["qvalue"]), _FLOOR))), fontsize=5)
ax.axhline(-np.log10(0.05), color="grey", ls="--", lw=1)
_count_either_side(ax, significance, -np.log10(0.05))
ax.axvline(0, color="black", lw=0.6)
ax.set_xlabel("effect size")
ax.set_ylabel("-log10 q")
ax.set_title(str(group), fontsize=9)
if palette:
ax.legend(
handles=[plt.Line2D([], [], color=colour, lw=6, label=label) for label, colour in palette.items()],
fontsize=5,
loc="upper left",
)
family_of = None if families is None else [str(families.get(name, "unknown")) for name in selected["feature"]]
tidy = pd.DataFrame(
{
"feature": selected["feature"].astype(str).to_numpy(),
"effect size": selected["effect"].to_numpy(dtype=float),
"-log10 q": significance,
}
)
if family_of is not None:
tidy["family"] = family_of
_maybe_interactive(
"scatter",
ax=ax,
data=tidy,
x="effect size",
y="-log10 q",
color="family" if family_of is not None else None,
hover=["feature"],
title=str(group),
)
return ax
[docs]
def dose_response(
adata: AnnData,
compound: str,
key: str = "dose_response",
compound_key: str = "Metadata_Compound",
dose_key: str = "Metadata_Concentration",
response: str | None = None,
ax: Axes | None = None,
) -> Axes:
"""One compound's response against dose, with the fitted curve when there is one.
Args:
adata: Object holding the table :func:`~mantispy.tl.dose_response` wrote.
compound: Which compound of that table to draw.
key: Name of that table in ``uns["mantispy"]``.
compound_key: ``obs`` column naming the compound of each well.
dose_key: ``obs`` column holding the dose of each well.
response: ``obs`` column drawn against the dose.
``None`` reads the column :func:`~mantispy.tl.dose_response` was given, so a plot cannot silently draw a different one than the table was fitted from.
ax: Axes to draw on, or ``None`` for a new figure.
Returns:
The axes drawn on, with a log ``y`` scale where every drawn response is positive, since a distance from the controls has a long right tail and one stray well would otherwise flatten the rest onto the baseline.
The scale is decided per axes, so panels drawn side by side can differ; set it on the returned axes to compare them.
Raises:
KeyError: There is no such table, or it holds no such compound.
"""
from mantispy.tl._dose import four_parameter_logistic
table = _table(adata, key, "mt.tl.dose_response")
row = _rows_for(table, "compound", compound, key)
if response is None:
response = _recorded_response(adata, "dose_response", "hits_row_distance")
obs = as_frame(adata.obs)
selected = obs[obs[compound_key].astype(str) == str(compound)]
doses = selected[dose_key].to_numpy(dtype=float)
values = selected[response].to_numpy(dtype=float)
usable = np.isfinite(doses) & np.isfinite(values) & (doses > 0)
ax = _axes(ax, (5, 4))
ax.scatter(doses[usable], values[usable], s=18, label="wells")
ax.set_xscale("log")
if usable.any() and values[usable].min() > 0:
ax.set_yscale("log")
fitted = row.iloc[0]
if bool(fitted["fit_ok"]):
grid = np.log10(np.geomspace(doses[usable].min(), doses[usable].max(), 100))
curve = four_parameter_logistic(
grid,
float(fitted["bottom"]),
float(fitted["top"]),
np.log10(float(fitted["ec50"])),
float(fitted["hill_slope"]),
)
ax.plot(10.0**grid, curve, color="crimson", lw=1.5, label=f"EC50 = {float(fitted['ec50']):.3g}")
dose_label = dose_key.replace("Metadata_", "")
ax.set_xlabel(dose_label)
ax.set_ylabel(response)
ax.set_title(f"{compound} (spearman {float(fitted['spearman']):.2f})", fontsize=9)
ax.legend(fontsize=7)
tidy = pd.DataFrame({dose_label: doses[usable], response: values[usable]})
if usable.any():
_maybe_interactive("scatter", ax=ax, data=tidy, x=dose_label, y=response, title=str(compound))
return ax
#: Grey where nothing happens, warm where it does, green where it has arrived, red where the cells are gone.
DOSE_PHASE_COLOURS = dict(zip(DOSE_PHASES, ("#f2f2f2", "#fde6c4", "#dbe8d4", "#f6d2d2"), strict=True))
[docs]
def dose_direction(
adata: AnnData,
compound: str,
key: str = "dose_direction",
ax: Axes | None = None,
) -> Axes:
"""One compound's ladder, with the background banded by what each concentration is doing.
The solid line is how far the profile sits from the controls, and the dashed line is how far it moved from the concentration below it.
The second is what says where the action is: a response that is still changing has a large step, one that has arrived has a small one however high the solid line sits.
The two dotted horizontals are the floors those lines are read against, which control wells laid out the same way reach.
Args:
adata: Object holding the table :func:`~mantispy.tl.dose_direction` wrote.
compound: Which compound of that table to draw.
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: There is no such table, or it holds no such compound.
"""
table = _table(adata, key, "mt.tl.dose_direction")
block = _rows_for(table, "compound", compound, key).sort_values("dose")
ax = _axes(ax, (5.2, 3.6))
doses = block["dose"].to_numpy(dtype=float)
# Each concentration's band reaches halfway to its neighbours in log dose, and half a step past the two ends.
log_dose = np.log10(doses)
gaps = np.diff(log_dose) if len(doses) > 1 else np.array([0.6])
padded = np.concatenate([[log_dose[0] - gaps[0]], log_dose, [log_dose[-1] + gaps[-1]]])
edges = 10.0 ** ((padded[:-1] + padded[1:]) / 2)
for left, right, phase in zip(edges[:-1], edges[1:], block["phase"], strict=True):
ax.axvspan(left, right, color=DOSE_PHASE_COLOURS.get(str(phase), "#ffffff"), lw=0, zorder=0)
floor = float(np.nanmedian(block["amplitude_null"].to_numpy(dtype=float)))
ax.axhline(floor, ls=":", lw=1, color="0.45", zorder=1)
ax.axhline(floor * np.sqrt(2.0), ls=":", lw=1, color="0.65", zorder=1)
ax.plot(
doses, block["amplitude"], marker="o", ms=4, lw=1.6, color="#2a4d69", label="distance from controls", zorder=3
)
ax.plot(
doses,
block["step_amplitude"],
marker="s",
ms=3,
lw=1.3,
ls="--",
color="#c1611f",
label="moved since the last",
zorder=3,
)
ax.set_xscale("log")
ax.set_xlim(edges[0], edges[-1])
ax.set_xlabel("concentration")
ax.set_ylabel("MADs per feature")
ax.set_title(str(compound), fontsize=10)
ax.legend(fontsize=7, frameon=False, loc="upper left")
tidy = pd.concat(
[
pd.DataFrame(
{
"concentration": doses,
"MADs per feature": block["amplitude"].to_numpy(dtype=float),
"series": "distance from controls",
}
),
pd.DataFrame(
{
"concentration": doses,
"MADs per feature": block["step_amplitude"].to_numpy(dtype=float),
"series": "moved since the last",
}
),
],
ignore_index=True,
)
_maybe_interactive(
"line", ax=ax, data=tidy, x="concentration", y="MADs per feature", color="series", title=str(compound)
)
return ax