Source code for mantispy.pp._transform

"""Rank-based inverse normal transformation.

This is the ``_int`` step of the JUMP profiling recipe for compound profiles and the third step of the baseline in :cite:t:`Arevalo_2024`.

Reference: the ``rank_int_array`` implementation in ``broadinstitute/jump-profiling-recipe``, which this reproduces on data with no missing values.
"""

from __future__ import annotations

import numpy as np
from anndata import AnnData

from mantispy._core._reduce import transform_grouped
from mantispy._core.logging import get_logger
from mantispy._core.mutation import inplace_or_copy

BLOM = 3.0 / 8.0


def rank_inverse_normal(X: np.ndarray, c: float = BLOM, stochastic: bool = True, seed: int = 0) -> np.ndarray:
    """Map each column onto a standard normal by rank.

    Args:
        X: Values to transform, one feature per column.
        c: Blom's constant in ``(rank - c) / (n - 2c + 1)``.
        stochastic: Break ties at random, as the reference implementation does.
            With ``False``, tied values share the mid-rank and get the same output, which suits real ties (such as a feature that is zero in half the wells) but not ties from rounding.
        seed: Seed for the tie-breaking.

    Returns:
        The transformed values, ``NaN`` where the input was missing.

    Notes:
        The finite values of each column are ranked among themselves, and only missing entries come back missing.
        The reference implementation passes the column to ``scipy.stats.rankdata``, which returns all NaN for a column with any missing value.
    """
    from scipy.special import ndtri
    from scipy.stats import rankdata

    values = np.atleast_2d(np.asarray(X, dtype=np.float64).T).T
    out = np.full(values.shape, np.nan)

    generator = np.random.default_rng(seed)
    order = generator.permutation(values.shape[0])

    for column in range(values.shape[1]):
        finite = np.flatnonzero(np.isfinite(values[:, column]))
        # One observation is still defined, ndtri((1 - c) / (2 - 2c)) = 0.0, so only an empty column is skipped.
        if finite.size == 0:
            continue
        present = values[finite, column]
        if stochastic:
            shuffle = order[np.isin(order, finite)]
            ranks = np.empty(finite.size)
            ranks[np.searchsorted(finite, shuffle)] = rankdata(values[shuffle, column], method="ordinal")
        else:
            ranks = rankdata(present, method="average")
        out[finite, column] = ndtri((ranks - c) / (finite.size - 2 * c + 1))

    return out.reshape(np.shape(X))


[docs] @inplace_or_copy() def rank_int( adata: AnnData, by: str | None = None, c: float = BLOM, stochastic: bool = True, seed: int = 0, key_added: str | None = None, copy: bool = False, ) -> AnnData | None: """Replace every feature by the normal quantile of its rank. Args: adata: Object to transform. Run it after :func:`~mantispy.pp.normalize`, as the JUMP recipe and the batch-correction benchmark do. by: Rank within each group of this column. ``None`` ranks globally, as the reference implementation does, which keeps every feature comparable across the screen. Ranking per plate also removes plate-level differences in distribution shape, but can hide a plate that failed. c: Blom's constant. stochastic: Tie handling; see ``rank_inverse_normal``. seed: Tie handling; see ``rank_inverse_normal``. key_added: Write to ``layers[key_added]`` instead of overwriting ``X``. copy: Return a modified copy instead of transforming in place. Returns: ``None``, or the modified copy. Writes ``X`` or ``layers[key_added]``. Notes: Every feature comes out standard normal, so no feature dominates a distance through its units. Effect sizes are lost as well: a feature that doubled and one that moved by one percent look the same if they reorder the same wells. Keep the untransformed values with ``key_added`` for effect sizes and dose-response curves. """ out = transform_grouped( adata, by, lambda _key, block: rank_inverse_normal(block, c=c, stochastic=stochastic, seed=seed) ) if key_added is None: adata.X = out else: adata.layers[key_added] = out get_logger().info("rank_int transformed %d features within %s", adata.n_vars, by or "the whole object") return None