Source code for mantispy.tl._ora

from __future__ import annotations

import numbers
from typing import Literal

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

from mantispy._core._stats import benjamini_hochberg
from mantispy._core.frames import as_frame
from mantispy._core.logging import get_logger
from mantispy._core.mutation import inplace_or_copy


[docs] @inplace_or_copy() def ora( adata: AnnData, groupby: str = "cluster", net: pd.DataFrame | None = None, gene_key: str = "Metadata_Gene", source: str = "source", target: str = "target", tmin: int = 5, key_added: str = "ora", copy: bool = False, *, padj_by: Literal["all", "group"] = "all", min_overlap: int = 1, ) -> AnnData | None: """Test each group's genes for over-representation of gene sets. Args: adata: Object whose ``obs`` carries the grouping and the gene symbols, for example after :func:`~mantispy.tl.cluster`. groupby: ``obs`` column defining the groups whose genes are tested. net: A gene-set network with ``source`` and ``target`` columns, such as :func:`mantispy.ds.gene_sets` returns. gene_key: ``obs`` column holding the gene symbol. source: Column of ``net`` naming the set. Renamed to ``source`` internally. target: Column of ``net`` naming the gene. Renamed to ``target`` internally. tmin: Smallest number of a set's genes that must be in the universe for the set to be tested. key_added: Name for the output table. copy: Return a modified copy instead of mutating in place. padj_by: Scope of the Benjamini-Hochberg correction. ``"all"`` (default) corrects once across every group and set. ``"group"`` corrects within each group's tests, so a group's modest enrichment is not penalized by unrelated groups (use this when many groups are tested at once). Because the scope is chosen per call, q-values from an ``"all"`` run and a ``"group"`` run are not directly comparable, so keep one scope within a single comparison. min_overlap: Smallest number of a group's genes a set must contain to be tested for that group. ``1`` (default) drops only the sets a group's genes do not hit at all (``a == 0``), so a set that shares no gene with the group is not tested and does not enlarge the correction denominator. The test is two-tailed, so a retained set may come out over- OR under-represented (read the sign of ``odds_ratio``); ``min_overlap`` bounds only how many of the group's genes a set must contain, not the direction of the result. Values ``>= 2`` require stronger overlap before a set is tested. Returns: ``None``, or the modified copy. Writes ``uns["mantispy"][key_added]`` with ``group``, ``source`` (the set), ``n`` (the group's genes in that set), ``odds_ratio`` (the Haldane-Anscombe log odds ratio), ``pvalue`` (a two-tailed Fisher exact test) and ``qvalue`` (Benjamini-Hochberg corrected), sorted by q. Only sets with at least ``min_overlap`` of a group's genes appear in that group's rows. Raises: KeyError: ``obs`` has no ``groupby`` or no ``gene_key``. ValueError: ``net`` is not given, or its ``source``/``target`` columns are missing. Notes: The universe is the set of distinct genes in ``obs[gene_key]``, so a set is tested only on its genes that the screen measured, and sets with fewer than ``tmin`` measured genes are skipped. Each group and set is tested with a two-tailed Fisher exact test over that universe. ``min_overlap`` shrinks the set of tested hypotheses to the sets a group's genes actually hit (overlap ``>= min_overlap``), rather than the whole collection. Sets a group does not touch are never tested, so they do not weigh on the correction under either scope: with ``padj_by="group"`` they stay out of that group's own Benjamini-Hochberg family, and with ``padj_by="all"`` (default) they are absent from the single pooled family shared across groups. """ from scipy.stats import fisher_exact if padj_by not in ("all", "group"): raise ValueError(f"padj_by must be 'all' or 'group', got {padj_by!r}") if isinstance(min_overlap, bool) or not isinstance(min_overlap, numbers.Integral) or min_overlap < 1: raise ValueError(f"min_overlap must be an int >= 1, got {min_overlap!r}") obs = as_frame(adata.obs) for column in (groupby, gene_key): if column not in obs: raise KeyError(f"obs has no column {column!r}") if net is None: raise ValueError("net is required; pass a gene-set network, e.g. mt.ds.gene_sets('hallmark')") network = net.rename(columns={source: "source", target: "target"}) if not {"source", "target"} <= set(network.columns): raise ValueError(f"net must have columns {source!r} and {target!r}") gene_names = obs[gene_key].astype(str).to_numpy() present = obs[gene_key].notna().to_numpy() & (gene_names != "") in_universe = set(gene_names[present]) n_bg = len(in_universe) network = network.astype({"source": str, "target": str}) network = network[network["target"].isin(in_universe)] set_genes = {name: set(block["target"]) for name, block in network.groupby("source", observed=True)} set_genes = {name: targets for name, targets in set_genes.items() if len(targets) >= tmin} group_labels = obs[groupby].astype(str).to_numpy() groups = pd.unique(group_labels) records = [] for group in groups: members = set(gene_names[(group_labels == group) & present]) & in_universe k = len(members) if k == 0 or k == n_bg: continue for name, targets in set_genes.items(): a = len(members & targets) if a < min_overlap: continue n_s = len(targets) # 2x2: rows member/non-member genes, columns in-set/out-of-set, over the measured universe. b, c = k - a, n_s - a d = n_bg - k - c records.append( { "group": group, "source": str(name), "n": a, # Haldane-Anscombe log odds ratio: +0.5 per cell keeps it finite when a cell is zero. "odds_ratio": float(np.log(((a + 0.5) * (d + 0.5)) / ((b + 0.5) * (c + 0.5)))), "pvalue": float(fisher_exact([[a, b], [c, d]], alternative="two-sided")[1]), } ) table = pd.DataFrame(records, columns=["group", "source", "n", "odds_ratio", "pvalue"]) if padj_by == "group": table["qvalue"] = table.groupby("group", observed=True, sort=False)["pvalue"].transform( lambda p: benjamini_hochberg(p.to_numpy()) ) else: table["qvalue"] = benjamini_hochberg(table["pvalue"].to_numpy()) table = table.sort_values("qvalue").reset_index(drop=True) adata.uns.setdefault("mantispy", {})[key_added] = table get_logger().info("ora: %d test(s) over %d group(s)", len(table), len(groups)) return None