Source code for mantispy.metrics._diagnose

from __future__ import annotations

import warnings
from collections.abc import Callable

import anndata as ad
import numpy as np
import pandas as pd
from anndata import AnnData

from mantispy._core._reduce import get_matrix, group_codes, group_offsets
from mantispy._core._stats import _default_well_block, benjamini_hochberg
from mantispy._core.frames import as_frame
from mantispy._core.masks import reference_mask
from mantispy._core.schema import RESOLUTIONS, stamp

#: How far the observed null rate may exceed the nominal one before it is a failure.
TOLERANCE = 2.0


def _verdict(ok: bool, warn: bool = False) -> str:
    return "pass" if ok else ("warn" if warn else "FAIL")


def _empirical_null(controls: AnnData, size: int, n_draws: int, seed: int, block: str | None) -> np.ndarray:
    """P-values from relabeling control wells as a treatment of ``size`` wells."""
    from mantispy.tl._differential import differential_features

    values = get_matrix(controls)
    pvalues = []
    for draw in range(n_draws):
        rng = np.random.default_rng(seed + draw)
        picked = rng.choice(controls.n_obs, size=size, replace=False)
        labels = np.full(controls.n_obs, "__reference__", dtype=object)
        labels[picked] = "__pseudo__"
        obs = as_frame(controls.obs).copy()
        obs["Metadata_Perturbation"] = labels
        obs["Metadata_Control"] = labels == "__reference__"
        scratch = ad.AnnData(X=values.copy(), obs=obs, var=as_frame(controls.var).copy())
        stamp(scratch, resolution="well")
        with warnings.catch_warnings():
            warnings.simplefilter("ignore")
            differential_features(scratch, block=block, key_added="__null__")
        table = scratch.uns["mantispy"]["__null__"]
        pvalues.append(table.loc[table["group"] == "__pseudo__", "pvalue"].to_numpy())
    return np.concatenate(pvalues) if pvalues else np.empty(0)


def _empirical_hit_rate(
    controls: AnnData, size: int, n_draws: int, seed: int, n_permutations: int, alpha: float
) -> dict[str, int]:
    """Count the control-only pseudo-treatments each hit caller calls."""
    from mantispy.tl._distance import edistance
    from mantispy.tl._hits import hit_calling

    values = get_matrix(controls)
    called = {"hit_calling": 0, "edistance": 0}
    for draw in range(n_draws):
        rng = np.random.default_rng(seed + draw)
        picked = rng.choice(controls.n_obs, size=size, replace=False)
        labels = np.where(np.isin(np.arange(controls.n_obs), picked), "__pseudo__", "__reference__")
        obs = as_frame(controls.obs).copy()
        obs["Metadata_Perturbation"] = labels
        obs["Metadata_Control"] = labels == "__reference__"
        scratch = ad.AnnData(X=values.copy(), obs=obs, var=as_frame(controls.var).copy())
        stamp(scratch, resolution="well")
        with warnings.catch_warnings():
            warnings.simplefilter("ignore")
            for name, function in (("hit_calling", hit_calling), ("edistance", edistance)):
                function(scratch, n_permutations=n_permutations, seed=draw, key_added=f"__{name}__")
                table = scratch.uns["mantispy"][f"__{name}__"]
                row = table.loc[table["group"] == "__pseudo__", "pvalue"]
                called[name] += int((row.to_numpy() < alpha).sum())
    return {name: int(value) for name, value in called.items()}


def _empirical_cell_hit_rate(
    controls: AnnData,
    well_codes: np.ndarray,
    pseudo_wells: int,
    n_draws: int,
    seed: int,
    n_permutations: int,
    alpha: float,
    *,
    method: str,
    well_block: str | None,
) -> dict[str, dict[str, int]]:
    """Count control-only pseudo-treatments each hit caller calls, under two nulls side by side.

    A pseudo-treatment is every cell of a random subset of control wells, kept whole, since the well
    structure is what makes the two nulls differ. Each caller runs twice per draw:

    - ``well-block``: the caller with the ``well_block`` unit (``None`` is the physical well), which draws
      whole wells for the null, the design's exchangeable unit. This is the calibrated rate.
    - ``cell-shuffle``: the same caller given a block that puts every cell in its own group, so its own
      well-block permutation engine resamples single cells. Cells within a well are not independent
      replicates, so this rate is anti-conservative and is the number to distrust.

    Both go through the caller's one code path, so the difference is the null unit alone.
    ``hit_calling`` runs the ``method`` the caller passed, ``"ks"`` at cell resolution by default.
    """
    from mantispy.tl._distance import edistance
    from mantispy.tl._hits import hit_calling

    # Each caller reuses one result key, with the extra keyword only hit_calling takes.
    callers: tuple[tuple[str, Callable[..., object], dict[str, str]], ...] = (
        ("hit_calling", hit_calling, {"method": method}),
        ("edistance", edistance, {}),
    )
    values = get_matrix(controls)
    wells = np.unique(well_codes)
    called: dict[str, dict[str, int]] = {name: {"well-block": 0, "cell-shuffle": 0} for name, _, _ in callers}
    for draw in range(n_draws):
        rng = np.random.default_rng(seed + draw)
        picked = rng.choice(wells, size=pseudo_wells, replace=False)
        labels = np.where(np.isin(well_codes, picked), "__pseudo__", "__reference__")
        obs = as_frame(controls.obs).copy()
        obs["Metadata_Perturbation"] = labels
        obs["Metadata_Control"] = labels == "__reference__"
        # One block per cell: whole-block resampling then draws single cells, a free cell shuffle.
        obs["__cell__"] = np.arange(controls.n_obs)
        scratch = ad.AnnData(X=values.copy(), obs=obs, var=as_frame(controls.var).copy())
        stamp(scratch, resolution="cell")
        with warnings.catch_warnings():
            warnings.simplefilter("ignore")
            for mode, block in (("well-block", well_block), ("cell-shuffle", "__cell__")):
                for name, function, extra in callers:
                    function(
                        scratch, block=block, n_permutations=n_permutations, seed=draw, key_added="__result__", **extra
                    )
                    table = scratch.uns["mantispy"]["__result__"]
                    row = table.loc[table["group"] == "__pseudo__", "pvalue"].to_numpy()
                    called[name][mode] += int((row < alpha).sum())
    return called


def _diagnose_cell(
    adata: AnnData,
    groupby: str,
    reference: str | None,
    block: str | None,
    n_draws: int,
    alpha: float,
    seed: int,
    n_permutations: int,
    method: str,
) -> pd.DataFrame:
    """Cell-resolution calibration: expose cell-within-well pseudoreplication.

    Real control cells are relabeled as pseudo-treatments of whole control wells, and the false positive
    rate of each hit caller is reported under the naive cell-shuffle null and the well-block null side by
    side. See :func:`_empirical_cell_hit_rate` and the ``Notes`` of :func:`diagnose_testing`.
    """
    from scipy import stats

    # At cell resolution `block` names the exchangeable unit for the permutation null, the well. The
    # well-level default "Metadata_Plate" (a blocking covariate there) names no such unit, so it is read as
    # the physical well; any other value is honored as the user's unit, matching the null their test runs.
    well_unit = None if block == "Metadata_Plate" else block
    block_name = "the physical well" if well_unit is None else f"the {well_unit!r} block"

    is_control = reference_mask(adata, reference)
    codes, keys = group_codes(adata, groupby)
    if well_unit is not None and well_unit not in adata.obs.columns:
        raise ValueError(
            f"diagnose_testing at cell resolution was asked to block on {well_unit!r}, which is not an obs "
            "column; pass a well column for the well-block null, or leave block at its default for the physical well"
        )
    well_codes = _default_well_block(adata, block=well_unit)
    if well_codes is None:
        raise ValueError(
            "diagnose_testing at cell resolution needs a complete (Metadata_Plate, Metadata_Well) with "
            "replicated wells to draw the well-block null; add it, or aggregate to wells with mt.tl.aggregate"
        )

    # Wells per treatment, the analogue of the well-level sizes, so the pseudo-treatment matches yours.
    treatment_wells = [
        int(np.unique(well_codes[codes == index]).size)
        for index in range(len(keys))
        if not is_control[codes == index].all()
    ]
    scored = [count for count in treatment_wells if count >= 2]
    controls = adata[is_control].copy()
    control_wells = _default_well_block(controls, block=well_unit)
    n_control_wells = 0 if control_wells is None else int(np.unique(control_wells).size)
    if not scored:
        raise ValueError("need at least one treatment on two wells to measure a null against")
    # A genuine well-block null needs four reference wells left after the pseudo-treatment is drawn, since
    # hit_calling keeps the well block only with four or more reference wells and otherwise falls back to a
    # per-cell KS test, plus a pseudo-treatment of at least two wells: at least six reference wells in all.
    if control_wells is None or n_control_wells < 6:
        raise ValueError(
            "diagnose_testing at cell resolution needs at least six reference wells to draw a genuine "
            "well-block null (four to form the null, two for the pseudo-treatment); "
            f"{block_name} gives {n_control_wells}. Aggregate to wells with mt.tl.aggregate to test at well level."
        )
    assert control_wells is not None  # narrowed by the guard above; the well-block call below needs it non-None

    typical = int(np.median(scored))
    # Cap the pseudo-treatment so at least four reference wells remain for a real well-block null, with the
    # two-well floor the design needs; the six-well guard above keeps the floor from ever breaking the cap.
    pseudo_wells = max(min(typical, n_control_wells - 4), 2)
    counts = _empirical_cell_hit_rate(
        controls, control_wells, pseudo_wells, n_draws, seed, n_permutations, alpha, method=method, well_block=well_unit
    )

    # A calibrated test calls a pseudo-treatment at rate alpha, so the count over n_draws is
    # Binomial(n_draws, alpha) and the cutoff is its 95th percentile, as the well-level hit-caller rows use.
    critical = int(stats.binom.ppf(0.95, n_draws, alpha))
    rows = []
    for name in ("hit_calling", "edistance"):
        well_block = counts[name]["well-block"]
        cell_shuffle = counts[name]["cell-shuffle"]
        measured = f"method={method!r} and {block_name}" if name == "hit_calling" else block_name
        rows.append(
            {
                "check": f"{name} well-block null rate",
                "value": f"{well_block} of {n_draws}",
                "expected": f"<= {critical}",
                "verdict": _verdict(well_block <= critical, warn=well_block <= critical + 1),
                "note": (
                    f"{name} called a control-only pseudo-treatment of {pseudo_wells} wells at p<{alpha} in "
                    f"{well_block} of {n_draws} draws, drawing whole wells for the null with {measured} as the cell "
                    f"design requires; a calibrated test exceeds {critical} about 5% of the time by chance. This is "
                    "the rate to trust."
                ),
            }
        )
        # Pseudoreplication is present when the cell-shuffle count exceeds both the well-block rate and the
        # chance cutoff; a clean screen where cells are roughly exchangeable keeps the two rates in line.
        inflated = cell_shuffle > well_block and cell_shuffle > critical
        rows.append(
            {
                "check": f"{name} cell-shuffle null rate",
                "value": f"{cell_shuffle} of {n_draws}",
                "expected": f"<= {well_block} (well-block)",
                "verdict": _verdict(not inflated, warn=inflated and cell_shuffle <= critical + 1),
                "note": (
                    f"the same {name} null permuting single cells instead of whole wells, called in {cell_shuffle} "
                    f"of {n_draws} draws. Cells within a well are not independent replicates, so a rate above both "
                    f"the well-block rate of {well_block} and the chance cutoff of {critical} is pseudoreplication"
                    + (
                        "; it is inflated here, so cells are not exchangeable on this screen."
                        if inflated
                        else "; it is in line with the well-block rate here, so cells are roughly exchangeable."
                    )
                ),
            }
        )
    return pd.DataFrame(rows, columns=["check", "value", "expected", "verdict", "note"])


[docs] def diagnose_testing( adata: AnnData, groupby: str = "Metadata_Perturbation", reference: str | None = "negcon", block: str | None = "Metadata_Plate", n_draws: int = 8, alpha: float = 0.05, seed: int = 0, n_permutations: int = 200, *, method: str = "ks", ) -> pd.DataFrame: """Check whether differential testing is calibrated on this screen. Args: adata: Well-level profiles after the normalization and transform you plan to test with, since the results depend on both. groupby: As in :func:`~mantispy.tl.differential_features`. reference: As in :func:`~mantispy.tl.differential_features`. block: At well resolution, as in :func:`~mantispy.tl.differential_features`. At cell resolution it names the exchangeable unit for the permutation null, normally the well; the well-level default ``"Metadata_Plate"`` is read there as the physical well, and any other column is honored so the diagnosed null equals the null your test would run. n_draws: Pseudo-treatments drawn from the controls for the empirical null. More draws resolve the false positive rate better and take longer. alpha: Nominal rate the null is compared against. seed: Seed for choosing which control wells stand in for a treatment. The two hit callers' permutation nulls are seeded by the draw index instead, so they are identical across calls that differ only in ``seed``. n_permutations: Null size for the two hit callers; smaller is faster and coarser. method: Keyword-only. The :func:`~mantispy.tl.hit_calling` method the cell-resolution checks measure, ``"ks"`` by default because the diagnostic is about pseudoreplication. ``"mahalanobis"`` whitens the distances, which divides out the between-well signal the cell-shuffle null exists to expose, so it cannot show that inflation. It is not used at well resolution, which keeps each caller's own default. Returns: A frame with columns ``check``, ``value``, ``expected``, ``verdict`` and ``note``, one row per check that ran, where a ``FAIL`` verdict means the check does not hold on this data. The empirical-null rows are absent when every null p-value came back non-finite, and the two hit-caller rows need at least eight reference wells. Raises: ValueError: At well resolution, no treatment has two wells, or there are fewer than four reference wells, leaving nothing to measure a null against. At cell resolution, the object is not stamped a known resolution (stamp it with :func:`mantispy.io.stamp`, or aggregate to wells with :func:`~mantispy.tl.aggregate`); or it has no complete, replicated well column, or a ``block`` naming no ``obs`` column, to draw the well-block null from; or fewer than six reference wells, too few to leave four for the null after a two-well pseudo-treatment. Notes: At cell resolution the checks change, because cells within a well are not independent replicates and the well-level checks describe well-level testing. Control cells are relabeled as pseudo-treatments of whole control wells and put through the two cell-resolution hit callers under two nulls. Each row's ``note`` names the block the null draws whole, and the ``hit_calling`` rows name ``method`` too, so there is no silent substitution of the test you meant: ``hit_calling well-block null rate`` / ``edistance well-block null rate`` The caller drawing whole wells for its null, the design's exchangeable unit. The count is compared against the upper tail of ``Binomial(n_draws, alpha)`` as the well-level hit-caller rows are. This is the rate to trust. ``hit_calling cell-shuffle null rate`` / ``edistance cell-shuffle null rate`` The same caller with the null permuting single cells instead of whole wells. Cells within a well share the well, so this shrinks the null spread by the cell count rather than the well count and runs above the well-block rate when cells are not exchangeable. The verdict compares the two: it fails when the cell-shuffle count exceeds both the well-block rate and the chance cutoff, the signature of pseudoreplication, and passes when the two rates agree. At well resolution the checks and what each one detects: ``null p < 0.05`` / ``null p < 0.01`` Control wells relabeled as treatments of the size yours have. The rate should match the nominal one. Heavy tails distort small p-values first, so a test can be calibrated at 0.05 and not at 0.01, which is closer to the range a false discovery rate works in. ``null discoveries`` How many of those null p-values survive Benjamini-Hochberg. A count above zero means the q-values on the real data are optimistic by roughly that much. ``hit_calling null rate`` / ``edistance null rate`` The same relabeling applied to the two hit callers, counted over ``n_draws`` draws. Both are permutation tests that are not fully calibrated at small control counts, so the count is compared against the upper tail of ``Binomial(n_draws, alpha)`` instead of a fixed rate. At eight draws the smallest non-zero rate is 0.125, and a threshold below that would fail a calibrated screen a third of the time. Raising ``n_draws`` sharpens the answer and moves the cutoff with it. ``rank test resolution`` The smallest p-value a Mann-Whitney test can return at your replication, compared with what multiple-testing correction requires. With three wells against 14 reference wells the floor is 2.9e-03 whatever the effect size, and :func:`~mantispy.tl.effect_size` then silently calls nothing. ``excess kurtosis`` How far the features are from the normality a t-test assumes. It predicts the null checks above but is not a verdict on its own, since heavy tails matter less with enough wells per group. ``wells per treatment`` and ``treatments sharing a {block} with the reference`` The replicate structure the other checks depend on. A treatment whose wells share no block with the reference cannot be tested. """ from scipy import stats resolution = adata.uns.get("mantispy", {}).get("resolution") if resolution == "cell": return _diagnose_cell(adata, groupby, reference, block, n_draws, alpha, seed, n_permutations, method) if resolution not in RESOLUTIONS: # get_resolution would default an unstamped object to "cell" and route it to the cell path, which then # cannot tell whether its wells are single cells; name the two fixes instead of a missing-column error. raise ValueError( "diagnose_testing cannot tell this object's resolution because it is not stamped, and the checks " "differ by resolution. Stamp it with mt.io.stamp(resolution=...), or aggregate single cells to " "wells with mt.tl.aggregate first." ) obs = as_frame(adata.obs) is_control = reference_mask(adata, reference) codes, keys = group_codes(adata, groupby) sizes = np.bincount(codes[~is_control], minlength=len(keys)) scored = sizes[sizes >= 2] if not scored.size or is_control.sum() < 4: raise ValueError("need at least one treatment with two wells and four reference wells") typical = int(np.median(scored)) n_tests = int(scored.size) * adata.n_vars threshold = alpha / max(n_tests, 1) rows = [] rows.append( { "check": "wells per treatment", "value": f"{typical} (min {int(scored.min())})", "expected": ">= 3", "verdict": _verdict(bool(scored.min() >= 3), warn=bool(scored.min() >= 2)), "note": "the unit that was randomized, and the sample size of every test", } ) if block is not None and block in obs.columns: blocks = obs[block].to_numpy() control_blocks = set(blocks[is_control]) order, offsets = group_offsets(codes, len(keys)) stranded: list[str] = [] spans: list[int] = [] for index in range(len(keys)): if sizes[index] < 2: continue group = order[offsets[index] : offsets[index + 1]] seen = set(blocks[group[~is_control[group]]]) spans.append(len(seen)) if not seen & control_blocks: stranded.append(str(keys[index])) rows.append( { "check": f"treatments sharing a {block} with the reference", "value": f"{len(spans) - len(stranded)} of {len(spans)}", "expected": "all", "verdict": _verdict(not stranded), "note": "a treatment on plates with no controls cannot be told from its plate" if stranded else f"median {int(np.median(spans))} {block} per treatment", } ) values = get_matrix(adata).astype(np.float64) finite = np.isfinite(values).all(axis=0) centred = values[:, finite] - values[:, finite].mean(axis=0) variance = np.mean(centred**2, axis=0) with np.errstate(invalid="ignore", divide="ignore"): kurtosis = float(np.nanmean(np.mean(centred**4, axis=0) / np.where(variance > 0, variance**2, np.nan) - 3.0)) rows.append( { "check": "excess kurtosis", "value": f"{kurtosis:.1f}", "expected": "0 (Gaussian)", "verdict": _verdict(kurtosis <= 20, warn=True), "note": "heavy tails break the small p-values first; mt.pp.rank_int removes them", } ) # Distinct values, because scipy falls back to the normal approximation on tied samples. n_control = int(is_control.sum()) floor = float( stats.mannwhitneyu(np.arange(typical, dtype=float), np.arange(typical, typical + n_control, dtype=float)).pvalue ) needed = int(np.ceil(floor * n_tests / alpha)) rows.append( { "check": "rank test resolution", "value": f"{floor:.1e}", "expected": f"< {threshold:.1e}", "verdict": _verdict(floor < threshold), "note": f"smallest p a Mann-Whitney can return with {typical} wells against {n_control} reference wells; " + ( f"the top of {n_tests:,} tests needs {threshold:.1e}" if floor < threshold else f"under BH nothing is called until {needed:,} of {n_tests:,} tests reach it at once" ), } ) controls = adata[is_control].copy() # Cap a pseudo-treatment at half the control wells so the rest can serve as the reference. pseudo_size = max(min(typical, controls.n_obs // 2), 2) null = _empirical_null(controls, pseudo_size, n_draws, seed, block) null = null[np.isfinite(null)] if null.size: for level in (0.05, 0.01): observed = float(np.mean(null < level)) rows.append( { "check": f"null p < {level}", "value": f"{observed * 100:.1f}%", "expected": f"{level * 100:.0f}%", "verdict": _verdict(observed <= level * TOLERANCE), "note": f"control wells relabeled as {pseudo_size}-well treatments, {n_draws} draws" + ("" if pseudo_size == typical else f" (capped from {typical}: only {controls.n_obs} controls)"), } ) discoveries = int((benjamini_hochberg(null) < alpha).sum()) rows.append( { "check": "null discoveries", "value": f"{discoveries} of {null.size:,}", "expected": "0", "verdict": _verdict(discoveries == 0), "note": "findings on data where there is nothing to find", } ) if controls.n_obs >= 8: counts = _empirical_hit_rate(controls, pseudo_size, n_draws, seed, n_permutations, alpha) critical = int(stats.binom.ppf(0.95, n_draws, alpha)) for name, count in counts.items(): rows.append( { "check": f"{name} null rate", "value": f"{count} of {n_draws}", "expected": f"<= {critical}", "verdict": _verdict(count <= critical, warn=count <= critical + 1), "note": ( f"{name} called a control-only pseudo-treatment of {pseudo_size} wells at " f"p<{alpha} in {count} of {n_draws} draws; a calibrated test exceeds " f"{critical} about 5% of the time by chance. Raise n_draws for a sharper " "answer; the cutoff moves with it." ), } ) return pd.DataFrame(rows, columns=["check", "value", "expected", "verdict", "note"])