Source code for mantispy.metrics._evaluate

"""The integration benchmark: scib-metrics' standard panel drawn as a mantispy heatmap, with a copairs mAP block when copairs is present."""

from __future__ import annotations

import warnings
from collections.abc import Sequence
from typing import TYPE_CHECKING, Any

import numpy as np
import pandas as pd

from mantispy._core.masks import reference_mask
from mantispy.metrics._common import embedding

if TYPE_CHECKING:
    from anndata import AnnData
    from matplotlib.axes import Axes

#: The copairs pair-argument keys a caller must fully specify to bypass a preset.
_PAIR_KEYS = ("pos_sameby", "pos_diffby", "neg_sameby", "neg_diffby")


def _has_scib() -> bool:
    """Whether scib-metrics, the engine behind the benchmark panel, can be imported."""
    try:
        import scib_metrics  # noqa: F401
    except ImportError:
        return False
    return True


def _has_copairs() -> bool:
    """Whether copairs, which adds the mean-average-precision block, can be imported."""
    try:
        from copairs import map as _  # noqa: F401
    except ImportError:
        return False
    return True


def _map_settings(
    map_mode: str, label_key: str, batch_key: str, map_kwargs: dict[str, Any] | None
) -> dict[str, list[str]]:
    """The four copairs pair arguments for the mAP, from the ``replicability``/``cross_plate`` preset plus overrides.

    Raises:
        ValueError: ``map_mode`` is neither preset and ``map_kwargs`` did not fully specify the four pair arguments.
    """
    from mantispy.tl._map import MODES, _resolve_mode

    if map_mode in MODES and map_mode not in ("activity", "consistency"):
        settings = _resolve_mode(map_mode, label=label_key, batch=batch_key)
        if map_kwargs is not None:
            settings.update({key: list(value) for key, value in map_kwargs.items()})
        return settings
    if map_kwargs is None or not set(_PAIR_KEYS) <= set(map_kwargs):
        raise ValueError(
            f"map_mode must be 'replicability' or 'cross_plate', or pass map_kwargs with all of {_PAIR_KEYS}; "
            f"got {map_mode!r}"
        )
    return {key: list(map_kwargs[key]) for key in _PAIR_KEYS}


def _map_meta(adata: AnnData, settings: dict[str, list[str]]) -> tuple[pd.DataFrame, np.ndarray | None]:
    """The copairs metadata frame and the treated-row mask, built once so every representation reuses them."""
    from mantispy._core.frames import as_frame

    obs = as_frame(adata.obs)
    needed = sorted({column for group in settings.values() for column in group})
    meta = pd.DataFrame({column: obs[column].reset_index(drop=True).to_numpy() for column in needed})
    treated = None
    if "Metadata_Control" in adata.obs.columns:
        # reference_mask is the hardened control read: it handles the "True"/"False" categorical an h5ad round trip
        # leaves and raises on a NaN control rather than silently dropping it, matching mt.tl.map.
        treated = ~reference_mask(adata, "negcon")
        meta = meta[treated].reset_index(drop=True)
    return meta, treated


def _rep_map(
    adata: AnnData,
    use_rep: str,
    settings: dict[str, list[str]],
    prepared: tuple[pd.DataFrame, np.ndarray | None] | None = None,
) -> float:
    """Mean of the per-group mean average precision over ``obsm[use_rep]``, scored with copairs, controls left out."""
    from mantispy.tl._map import _score_map

    meta, treated = prepared if prepared is not None else _map_meta(adata, settings)
    features = embedding(adata, use_rep).astype(np.float32)
    if treated is not None:
        features = features[treated]

    # compute_null=False skips the permutation null: only the mean point estimate is read here, not the p-values.
    table, _ = _score_map(meta, features, settings, null_size=0, threshold=0.05, seed=0, warn=False, compute_null=False)
    return float(table["mean_average_precision"].mean())


def _benchmark(
    adata: AnnData, *, reps: Sequence[str], label_key: str, batch_key: str, min_max_scale: bool
) -> pd.DataFrame:
    """scib-metrics' public ``get_results`` frame for the representations, kept as a seam a test can read directly.

    scib-metrics is imported here and only here, and only the public ``get_results`` is read, never the private
    ``_results``, so scib's own Total stays the number scib computes.
    """
    from scib_metrics.benchmark import Benchmarker

    # progress_bar=False keeps the tqdm bars out of executed notebooks and test logs.
    bm = Benchmarker(
        adata,
        batch_key=batch_key,
        label_key=label_key,
        embedding_obsm_keys=list(reps),
        n_jobs=-1,
        progress_bar=False,
    )
    bm.benchmark()
    return bm.get_results(min_max_scale=min_max_scale)


[docs] def evaluate_integration( adata: AnnData, *, reps: Sequence[str] = ("X_pca",), label_key: str = "Metadata_Perturbation", batch_key: str = "Metadata_Batch", min_max_scale: bool = False, map_mode: str = "replicability", map_kwargs: dict[str, Any] | None = None, ax: Axes | None = None, ) -> pd.DataFrame: """Score one or more representations against a batch and draw the integration benchmark as a heatmap. The numbers come from scib-metrics, the field's implementation of the integration panel, which mantispy reads through its public ``get_results`` rather than reimplements. mantispy renders its own heatmap from them and, when copairs is installed, adds the cross-replicate mean average precision (mAP) as its own block. The columns are grouped left to right into bio conservation, batch correction, retrieval (mAP) and the aggregate scores, each block under its own header. The bio-conservation metrics measure whether a representation keeps the biology (cLISI is how label-pure a well's neighbourhood is); the batch-correction metrics whether it mixes the batches (iLISI is how batch-mixed that neighbourhood is, PCR comparison how much less of the variance the batch explains after correction). ``Total`` is scib's own weighted score, ``0.4`` batch correction and ``0.6`` bio conservation, left exactly as scib computes it. ``Total+mAP`` is mantispy's: it folds mAP into the bio group as one more bio signal, ``bio' = mean(bio metrics + mAP)``, then reweights ``0.4`` batch correction and ``0.6`` bio'. The mAP is measured over the treated wells only (controls left out), while scib's bio, batch and ``Total`` columns use every well, so ``Total`` and ``Total+mAP`` are not strictly apples-to-apples on a control-heavy screen. The return type is uniform across install states, so the numbers are always reachable: - With ``mantispy[integration]`` (scib-metrics) and copairs, the frame holds every metric, the mAP column and both totals; the heatmap has all four blocks. - With only scib-metrics, the frame and heatmap drop the retrieval block and ``Total+mAP``. - With only copairs, the frame and heatmap hold the per-representation mAP alone, and a warning notes that scib-metrics adds the full panel. - With neither, this raises :class:`ImportError`. This function computes no native PC-regression of its own. To audit what a representation spends its variance on besides the label, such as the cell count or the plate position, whose better direction is context-dependent and so is deliberately not in this table, use :func:`~mantispy.metrics.pc_regression` for a single covariate or :func:`~mantispy.metrics.batch_variance_explained` to stack several. Args: adata: Object holding the representations in ``obsm`` and the label and batch in ``obs``. reps: ``obsm`` keys to compare, e.g. ``("X_pca", "X_pca_harmony")``. label_key: ``obs`` column with the biological grouping. batch_key: ``obs`` column with the nuisance grouping. min_max_scale: Colour and score each metric column scaled across the representations, as scib does by default. Off by default so a single representation still has honest absolute values to colour by. map_mode: The copairs pairing for mAP, ``"replicability"`` (the Arevalo ``mAP-nonrep`` default) or ``"cross_plate"``. See ``_map_settings``. map_kwargs: Overrides for any of the four copairs pair arguments, on top of ``map_mode``. ax: Axes to draw the heatmap on, or ``None`` for a new figure. Returns: The numeric results, one row per representation and one column per metric, the mAP and both totals. Raises: ImportError: Neither scib-metrics nor copairs is installed. ValueError: The object has one row per perturbation, which has no replicate pairs to score. KeyError: ``obs`` is missing ``label_key`` or ``batch_key``. """ from mantispy.pl._evaluation import _integration_heatmap have_scib = _has_scib() have_copairs = _has_copairs() if not have_scib and not have_copairs: raise ImportError( 'evaluate_integration needs scib-metrics and copairs; install pip install "mantispy[integration,map]". ' "For a native batch-variance check without them, use mt.metrics.batch_variance_explained (stack " "several covariates) or mt.metrics.pc_regression (a single covariate)." ) reps = tuple(dict.fromkeys(reps)) # a duplicate rep would draw two identically-labelled rows missing = [key for key in (label_key, batch_key) if key not in adata.obs.columns] if missing: raise KeyError(f"obs is missing the column(s) the benchmark needs: {missing}") # Count the treated labels, since the mAP drops the controls; value_counts drops NaN labels too. treated_labels = adata.obs[label_key] if "Metadata_Control" in adata.obs.columns: treated_labels = treated_labels[~reference_mask(adata, "negcon")] counts = treated_labels.value_counts() if counts.empty: raise ValueError( f"evaluate_integration found no treated profiles to score: obs[{label_key!r}] is empty or all missing " "once the controls are dropped." ) if int(counts.max()) < 2: raise ValueError( "evaluate_integration needs replicate profiles across batches, and this object looks like one row " f"per perturbation (no value of obs[{label_key!r}] has two or more profiles). It scores the " "integration of well-level replicates; for consensus profiles use mt.tl.map(mode='consistency') or " "mt.metrics.known_relationships instead." ) if have_copairs: settings = _map_settings(map_mode, label_key, batch_key, map_kwargs) prepared = _map_meta(adata, settings) map_values = [_rep_map(adata, rep, settings, prepared) for rep in reps] if not have_scib: warnings.warn( "scib-metrics is not installed, so evaluate_integration draws only the per-representation mean " 'average precision. Install pip install "mantispy[integration,map]" to add the full scib-metrics ' "panel (bio conservation, batch correction and the Total).", UserWarning, stacklevel=2, ) frame = pd.DataFrame( {"mean_average_precision": map_values}, index=pd.Index(list(reps), name="representation"), ) _integration_heatmap(frame, {"Retrieval": ["mean_average_precision"]}, ax=ax) return frame if min_max_scale and len(reps) == 1: warnings.warn( "min_max_scale scales each metric across the representations, and with a single representation there " "is nothing to scale against, so every scaled column collapses. Pass more than one representation, " "or keep the default min_max_scale=False for honest absolute values.", UserWarning, stacklevel=2, ) results = _benchmark(adata, reps=reps, label_key=label_key, batch_key=batch_key, min_max_scale=min_max_scale) metric_type = results.loc["Metric Type"] data = results.drop(index="Metric Type").astype(float).loc[list(reps)] bio_cols = [column for column in data.columns if metric_type[column] == "Bio conservation"] batch_cols = [column for column in data.columns if metric_type[column] == "Batch correction"] aggregate_cols = ["Batch correction", "Bio conservation", "Total"] blocks: dict[str, list[str]] = {"Bio conservation": bio_cols, "Batch correction": batch_cols} if have_copairs: map_column = pd.Series(map_values, index=data.index) if min_max_scale: # Fold the mAP in on the same worst->0, best->1 scale as the bio metrics it joins, so Total+mAP and # the heatmap colour are not a scaled/raw mix. span = map_column.max() - map_column.min() map_column = (map_column - map_column.min()) / span data["mean_average_precision"] = map_column blocks["Retrieval"] = ["mean_average_precision"] if bio_cols: bio_prime = (data[bio_cols].sum(axis=1) + data["mean_average_precision"]) / (len(bio_cols) + 1) data["Total+mAP"] = 0.4 * data["Batch correction"] + 0.6 * bio_prime aggregate_cols = [*aggregate_cols, "Total+mAP"] else: warnings.warn( "scib-metrics returned no bio-conservation columns, so Total+mAP is undefined and is left out; the " "mAP is still shown on its own. This usually means a scib-metrics version change.", UserWarning, stacklevel=2, ) blocks["Aggregate"] = aggregate_cols data = data[[column for columns in blocks.values() for column in columns]] _integration_heatmap(data, blocks, ax=ax) return data