Source code for campmc.mt.recovery

"""Rare cell-type recovery scores for metacell partitions."""

from __future__ import annotations

import pandas as pd
from anndata import AnnData
from sklearn.metrics import adjusted_rand_score, normalized_mutual_info_score


def _identify_rare_types(labels: pd.Series, *, rare_threshold: float = 0.01) -> pd.DataFrame:
    """Identify cell types whose frequency is at most ``rare_threshold``."""
    counts = labels.value_counts(dropna=False)
    freq = counts / counts.sum()
    rare = freq[freq <= rare_threshold].sort_values()
    return pd.DataFrame(
        {
            "cell_type": rare.index.astype(str),
            "n_cells": counts.loc[rare.index].values.astype(int),
            "fraction": rare.values.astype(float),
        }
    )


def _balanced_accuracy(true: pd.Series, pred: pd.Series, rare_types: list[str]) -> float:
    recalls: list[float] = []
    for ct in rare_types:
        pos_mask = true == ct
        n_pos = int(pos_mask.sum())
        if n_pos == 0:
            continue
        tp = int(((pred == ct) & pos_mask).sum())
        recalls.append(tp / n_pos)
    if not recalls:
        return float("nan")
    return float(sum(recalls) / len(recalls))


[docs] def recovery_scores( adata: AnnData, *, partition_key: str, label_key: str, rare_threshold: float = 0.01, ) -> pd.DataFrame: """Score how well a partition recovers rare cell types. Assigns each metacell the majority label among its members, then computes balanced accuracy over rare types plus ARI/NMI restricted to rare cells. Rare types are those with frequency at most ``rare_threshold``. Parameters ---------- adata Annotated data with partition and cell-type columns in ``obs``. partition_key Column in ``adata.obs`` with metacell assignments. label_key Column in ``adata.obs`` with ground-truth (or reference) labels. rare_threshold Maximum fraction of cells for a type to be considered rare. Returns ------- DataFrame Single-row table with ``balanced_accuracy``, ``ari``, ``nmi``, and ``n_rare_types``. Examples -------- .. exec-jupyter:: import campmc as cp import scanpy as sc adata = sc.datasets.pbmc68k_reduced() cp.partition(adata, method="camp3", gamma=50, random_state=0) cp.mt.recovery_scores( adata, partition_key="camp", label_key="louvain", rare_threshold=0.05, ) """ if partition_key not in adata.obs: raise KeyError(partition_key) if label_key not in adata.obs: raise KeyError(label_key) labels = adata.obs[label_key].astype(str) assign = adata.obs[partition_key].astype(str) rare_df = _identify_rare_types(labels, rare_threshold=rare_threshold) rare_types = rare_df["cell_type"].tolist() contingency = pd.crosstab(assign, labels) mc_majority = contingency.idxmax(axis=1) pred = assign.map(mc_majority).astype(str) mask = labels.isin(rare_types) ari = float("nan") nmi = float("nan") if mask.sum() > 1: ari = float(adjusted_rand_score(labels.loc[mask], pred.loc[mask])) nmi = float(normalized_mutual_info_score(labels.loc[mask], pred.loc[mask])) return pd.DataFrame( [ { "balanced_accuracy": _balanced_accuracy(labels, pred, rare_types), "ari": ari, "nmi": nmi, "n_rare_types": len(rare_types), } ] )