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),
}
]
)