"""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