Source code for mantispy.tl._consensus

"""Consensus signatures, one profile per perturbation.

The modz weighting follows ``pycytominer.cyto_utils.modz.modz_base``, which comes from cmapPy.
"""

from __future__ import annotations

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

from mantispy._core._numba import MEDIAN
from mantispy._core._reduce import get_matrix, group_codes, group_offsets, reduce_grouped, reduced_var
from mantispy._core.logging import report_drop
from mantispy._core.provenance import record_params
from mantispy._core.schema import stamp
from mantispy.tl._aggregate import _group_obs
from mantispy.tl._similarity import similarity_matrix

METHODS = ("modz", "median")
CORRELATIONS = ("spearman", "pearson")


def modz_weights(
    block: np.ndarray, correlation: str = "spearman", min_weight: float = 0.01, precision: int = 4
) -> np.ndarray:
    """Weight each replicate by how well it agrees with the others.

    A perturbation whose replicates all sit at ``min_weight`` has no reproducible signature, whatever its consensus profile looks like.

    Args:
        block: One perturbation's replicates, as rows, by features.
        correlation: ``"spearman"`` ranks the features first, ``"pearson"`` correlates the values.
        min_weight: Floor on a replicate's weight.
        precision: Decimals the weights are rounded to, as in pycytominer.

    Returns:
        One weight per row of ``block``, summing to one.
    """
    from scipy.stats import rankdata

    if block.shape[0] == 1:
        return np.ones(1)

    values = np.asarray(block, dtype=np.float64)
    if correlation == "spearman":
        # The default nan_policy, "propagate", turns a row with one missing feature into all NaN.
        values = rankdata(values, axis=1, nan_policy="omit")

    # Not similarity_matrix's zero fill: zero sits below every rank, so replicates sharing a gap would correlate.
    gaps = np.isnan(values)
    if gaps.any():
        present = np.maximum((~gaps).sum(axis=1, keepdims=True), 1)
        centre = np.nansum(values, axis=1, keepdims=True) / present
        values = np.where(gaps, centre, values)

    matrix = similarity_matrix(values, metric="pearson").astype(np.float64)
    np.fill_diagonal(matrix, np.nan)

    weights = np.nanmean(np.clip(matrix, 0.0, None), axis=1)
    weights = np.clip(weights, min_weight, None)
    total = weights.sum()
    weights = np.full(weights.size, 1.0 / weights.size) if total == 0 else weights / total
    return np.round(weights, precision)


[docs] def consensus( adata: AnnData, by: str = "Metadata_Perturbation", method: str = "median", correlation: str = "spearman", min_replicates: int = 2, min_weight: float = 0.01, precision: int = 4, use_rep: str | None = None, ) -> AnnData: """One profile per group, weighting replicates by how well they agree. Args: adata: Profiles to summarize, normally well level. by: Column defining a perturbation. method: ``"median"`` (the default, matching pycytominer and identical to :func:`~mantispy.tl.aggregate` by the same column) or ``"modz"``, which weights replicates by their agreement so a single bad replicate moves the signature far less than a plain mean would. correlation: How replicate agreement is measured: ``"spearman"`` (pycytominer's default, and insensitive to a few extreme features) or ``"pearson"``. min_replicates: Groups with fewer replicates are dropped. min_weight: Floor on a replicate's weight. A group whose replicates all land on the floor becomes an unweighted mean. precision: Decimals the weights are rounded to, as in pycytominer. use_rep: Reduce this ``obsm`` representation (e.g. an embedding from :func:`~mantispy.pp.tvn`/:func:`~mantispy.pp.harmony`) instead of ``X``; the result's ``X`` holds the reduced representation and ``var`` is a plain range index, since the axes are not named features. Returns: A new object at ``"perturbation"`` resolution, one row per group, with ``Metadata_ReplicateCount`` and the metadata that is constant within a group. ``uns["mantispy"]["consensus_weights"]`` keeps the weight given to every input row, including the rows of groups dropped for having too few replicates, so a signature can be traced back to its replicates. Under ``method="median"`` no weights are computed and every row is recorded as 1.0, since a median is not a weighted sum. Raises: ValueError: ``method`` is not one of ``METHODS``, ``correlation`` is not one of ``CORRELATIONS``, or ``use_rep`` is not a 2-D representation in ``obsm``. Notes: A missing value is filled with its own replicate's mean before the replicates are correlated. Zero would be an extreme value among ranks, and two replicates sharing a gap would look alike. The signature itself is a weighted sum, so a NaN feature stays NaN. modz is a weighted mean: with one outlying replicate it drifts about forty times less than the unweighted mean, but it does not beat a median. median is at least as robust as modz for consensus signatures, so compare both on your own data. Normalize before taking a consensus, and first drop the features :func:`~mantispy.pp.normalize` flags in ``var["degenerate_scale"]``. A feature that is constant among the controls is divided by epsilon, and a weighted mean carries the resulting values of order 1e17 into the signature, where a median would discard them. """ if method not in METHODS: raise ValueError(f"method must be one of {METHODS}, got {method!r}") if correlation not in CORRELATIONS: raise ValueError(f"correlation must be one of {CORRELATIONS}, got {correlation!r}") codes, keys = group_codes(adata, [by]) counts = np.bincount(codes, minlength=len(keys)) weights = np.ones(adata.n_obs) if method == "median": values, _, _ = reduce_grouped(adata, [by], MEDIAN, use_rep=use_rep) else: X = get_matrix(adata, use_rep=use_rep) values = np.zeros((len(keys), X.shape[1]), dtype=np.float64) order, offsets = group_offsets(codes, len(keys)) for index in range(len(keys)): rows = order[offsets[index] : offsets[index + 1]] block = modz_weights(X[rows], correlation, min_weight, precision) weights[rows] = block values[index] = block @ X[rows] obs = _group_obs(adata, [by], keys, codes, {"Metadata_ReplicateCount": counts}) keep = counts >= min_replicates report_drop("group(s)", int((~keep).sum()), int(keep.size), remedy=f"lower min_replicates below {min_replicates}") var = reduced_var(adata, use_rep, values.shape[1]) result = ad.AnnData( X=values[keep].astype(np.float32), obs=obs.loc[keep].reset_index(drop=True).set_axis(pd.Index([str(i) for i in range(int(keep.sum()))])), var=var, ) stamp(result, resolution="perturbation") result.uns["mantispy"]["consensus_weights"] = pd.DataFrame( {"group": [str(keys[code]) for code in codes], "weight": weights} ) record_params( result, "consensus", { "by": by, "method": method, "correlation": correlation, "min_replicates": min_replicates, "use_rep": use_rep, }, ) return result