Skip to content
50 changes: 50 additions & 0 deletions docs/api/tools_index.md
Original file line number Diff line number Diff line change
Expand Up @@ -295,6 +295,7 @@ Enrichment tests for single-cell data assess whether specific biological pathway
aiding in the identification of functional characteristics and cellular states.
While pathway enrichment is a well-studied and commonly applied approach in single-cell RNA-seq, other data sources such as genes targeted by drugs can also be enriched.
Drug2cell performs such enrichment tests and is available in pertpy {cite}`Kanemaru2023`.
The same enrichment interface can also score CMap-style signature reversal on perturbation-level data, ranking perturbations that most strongly oppose a query signature.

```{eval-rst}
.. autosummary::
Expand All @@ -315,6 +316,55 @@ pt_enricher = pt.tl.Enrichment()
pt_enricher.score(adata)
```

#### Signature reversal

`Enrichment.signature_reversal` computes a raw [weighted connectivity score (WTCS)](https://clue.io/connectopedia/cmap_algorithms) and stores both connectivity and its negative, the reversal score, in `adata.obs`.
Higher reversal scores indicate stronger transcriptional opposition to the query.
This is not a normalized CMap score: it does not compute NCS, tau, p-values, or false-discovery rates.

The input matrix must contain finite, signed perturbation effects relative to appropriate matched controls, such as z-scores, log-fold changes, or control-subtracted expression.
Do not use raw or pseudobulk mean expression directly.
The up and down sets should represent genes differentially expressed in the query state relative to its reference.
The following example aggregates the cell-level `distance_example()` data and subtracts its control profile:

```python
cell_adata = pt.dt.distance_example()
ps = pt.tl.PseudobulkSpace()
ps_adata = ps.compute(
cell_adata,
target_col="perturbation",
mode="mean",
)
ps_adata = ps.compute_control_diff(
ps_adata,
target_col="perturbation",
reference_key="control",
)
ps_adata = ps_adata[ps_adata.obs["perturbation"] != "control"].copy()

query_profile = -ps_adata[
ps_adata.obs["perturbation"] == "p-sgCREB1-2"
].to_df().iloc[0]
up_genes = query_profile[query_profile > 0].nlargest(20).index.tolist()
down_genes = query_profile[query_profile < 0].nsmallest(20).index.tolist()

enr = pt.tl.Enrichment()
enr.signature_reversal(
ps_adata,
up_genes=up_genes,
down_genes=down_genes,
)
```

To keep this example self-contained, its query is the opposite of one observed CRISPR perturbation signature.
In a real analysis, use an independently derived disease or state signature.
The [CMap query guidance](https://clue.io/connectopedia/how_to_construct_cmap_queries) recommends approximately 10 to 200 genes per query.
A signed query can also be supplied, but only the sign of each value determines whether a gene belongs to the up or down set.

A high reversal score is a hypothesis for follow-up, not evidence of therapeutic efficacy or safety.
In particular, a perturbation can score highly by suppressing a compensatory or protective stress response.
Results should therefore be interpreted together with biological context and orthogonal phenotypic, viability, and toxicity measurements.

See [enrichment tutorial](https://pertpy.readthedocs.io/en/latest/tutorials/notebooks/enrichment.html).

## Distances and permutation tests
Expand Down
226 changes: 226 additions & 0 deletions src/pertpy/tools/_enrichment.py
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@
import pandas as pd
import scanpy as sc
from anndata import AnnData
from fast_array_utils.conv import to_dense
from matplotlib.axes import Axes
from scanpy.plotting import DotPlot
from scanpy.tools._score_genes import _sparse_nanmean
Expand Down Expand Up @@ -56,6 +57,121 @@ def _mean(X, names, axis):
return obs_avg


def _get_signature_matrix(adata: AnnData, layer: str | None) -> np.ndarray | CSBase:
if layer is not None:
if layer not in adata.layers:
raise ValueError(f"Layer {layer!r} does not exist in the .layers attribute.")
matrix = adata.layers[layer]
else:
matrix = adata.X
return cast_matrix(matrix)


def _get_signature_genes(adata: AnnData, gene_symbols_key: str | None) -> np.ndarray:
raw_genes: pd.Index[Any] | pd.Series[Any]
if gene_symbols_key is None:
raw_genes = adata.var_names
else:
var = cast_frame(adata.var)
if gene_symbols_key not in var:
raise ValueError(f"Column {gene_symbols_key!r} does not exist in the .var attribute.")
raw_genes = var[gene_symbols_key]

if pd.isna(raw_genes).any():
raise ValueError("Gene identifiers used for signature matching must not be missing.")
genes = pd.Index(raw_genes.astype(str))
if genes.has_duplicates:
duplicates = genes[genes.duplicated()].unique().tolist()
raise ValueError(f"Gene identifiers used for signature matching must be unique; found {duplicates!r}.")
return genes.to_numpy()


def _prepare_query_signature(
query_signature: Mapping[str, float] | pd.Series | None,
up_genes: Sequence[str] | None,
down_genes: Sequence[str] | None,
) -> pd.Series:
if query_signature is not None and (up_genes is not None or down_genes is not None):
raise ValueError("Pass either `query_signature` or `up_genes`/`down_genes`, not both.")

if query_signature is not None:
if isinstance(query_signature, pd.Series):
signature = query_signature.astype(float).copy()
else:
signature = pd.Series(dict(query_signature), dtype=float)
else:
up = [up_genes] if isinstance(up_genes, str) else ([] if up_genes is None else list(up_genes))
down = [down_genes] if isinstance(down_genes, str) else ([] if down_genes is None else list(down_genes))
up = [str(gene) for gene in up]
down = [str(gene) for gene in down]
if len(up) != len(set(up)):
raise ValueError("`up_genes` must not contain duplicate genes.")
if len(down) != len(set(down)):
raise ValueError("`down_genes` must not contain duplicate genes.")

values: dict[str, float] = {}
for gene in up:
values[gene] = 1.0
for gene in down:
if gene in values:
raise ValueError(f"Gene {gene!r} occurs in both `up_genes` and `down_genes`.")
values[gene] = -1.0
signature = pd.Series(values, dtype=float)

if pd.isna(signature.index).any():
raise ValueError("Query gene identifiers must not be missing.")
signature.index = signature.index.astype(str)
if signature.index.has_duplicates:
duplicates = signature.index[signature.index.duplicated()].unique().tolist()
raise ValueError(f"Query gene identifiers must be unique; found {duplicates!r}.")
if not np.isfinite(signature.to_numpy(dtype=float)).all():
raise ValueError("Query signature values must be finite.")
signature = signature[signature != 0]
if signature.empty:
raise ValueError("The query signature must contain at least one non-zero gene.")
return signature


def _weighted_enrichment_score(values: np.ndarray, hits: np.ndarray) -> float:
n_hits = int(hits.sum())
n_misses = len(hits) - n_hits
if n_hits == 0 or n_misses == 0:
raise ValueError("Weighted enrichment requires at least one hit and one non-hit gene.")

order = np.argsort(values, kind="mergesort")[::-1]
ranked_values = values[order]
ranked_hits = hits[order]
ranked_weights = np.abs(ranked_values)
hit_weights = ranked_weights * ranked_hits
max_hit_weight = hit_weights.max()
if max_hit_weight == 0:
return 0.0

hit_weights = hit_weights / max_hit_weight
hit_weight_sum = hit_weights.sum()
hit_step = hit_weights / hit_weight_sum
miss_step = (~ranked_hits) / n_misses
tie_starts = np.r_[0, np.flatnonzero(ranked_values[1:] != ranked_values[:-1]) + 1]
running: np.ndarray = np.cumsum(np.add.reduceat(hit_step - miss_step, tie_starts))
running[-1] = 0.0
max_score = float(running.max())
min_score = float(running.min())
return max_score if abs(max_score) >= abs(min_score) else min_score


def _cmap_connectivity(values: np.ndarray, up_mask: np.ndarray, down_mask: np.ndarray) -> float:
es_up = _weighted_enrichment_score(values, up_mask) if up_mask.any() else float("nan")
es_down = _weighted_enrichment_score(values, down_mask) if down_mask.any() else float("nan")

if up_mask.any() and down_mask.any():
if np.sign(es_up) == np.sign(es_down):
return 0.0
return float((es_up - es_down) / 2.0)
if up_mask.any():
return es_up
return float(-es_down)


class Enrichment:
def score(
self,
Expand Down Expand Up @@ -152,6 +268,116 @@ def score(
adata.uns[f"{key_added}_genes"]["var"].loc[drug, "genes"] = "|".join(adata.var_names[target_groups[drug]])
adata.uns[f"{key_added}_all_genes"]["var"].loc[drug, "all_genes"] = "|".join(full_targets[drug])

def signature_reversal(
self,
adata: AnnData,
query_signature: Mapping[str, float] | pd.Series | None = None,
*,
up_genes: Sequence[str] | None = None,
down_genes: Sequence[str] | None = None,
layer: str | None = None,
gene_symbols_key: str | None = None,
min_genes: int = 1,
key_added: str = "signature_reversal",
) -> None:
"""Score perturbations by how strongly they reverse a disease/query signature.

This computes a raw CMap-style weighted connectivity score (WTCS) on a perturbation-level AnnData object, where observations are perturbations and variables are genes.
Values must be finite, signed perturbation effects relative to an appropriate matched control, such as z-scores, log-fold changes, or control-subtracted expression.
Raw or pseudobulk mean expression is not a perturbation signature.

Positive query genes are up-regulated in the query state and negative genes are down-regulated.
Only the sign of values in `query_signature` is used.
Connectivity ranges from -1 (opposing) to 1 (similar), and the stored reversal score is its negative so that higher values indicate stronger opposition.
This method does not compute CMap's normalized connectivity score, tau, p-values, or false-discovery rates.
Genes with equal perturbation values are treated as a single rank group so their input order cannot affect the score.

A high reversal score is a hypothesis for follow-up, not evidence of efficacy or safety.
In particular, suppressing a compensatory or protective transcriptional response can also produce a high reversal score.

Args:
adata: Perturbation-level AnnData with perturbations as observations, genes as variables, and finite signed perturbation effects relative to matched controls in `.X` or `layer`.
query_signature: Signed query signature.
Positive values indicate query-up genes and negative values query-down genes; magnitudes are ignored.
up_genes: Query-up genes. Used when `query_signature` is not provided.
down_genes: Query-down genes. Used when `query_signature` is not provided.
layer: Layer containing perturbation signatures. Defaults to `.X`.
gene_symbols_key: Optional `.var` column used to match query gene names instead of `.var_names`.
Gene identifiers used for matching must be unique and non-missing.
min_genes: Minimum total number of query genes that must be present in `adata`.
CMap recommends query sets containing roughly 10 to 200 genes; very small matches should be treated as exploratory.
key_added: Prefix used to store results in `.obs` and `.uns`.

Returns:
Updates `adata` with `{key_added}_score`, `{key_added}_connectivity` and `{key_added}_rank` in `.obs`, and query metadata in `.uns[key_added]`.

Examples:
>>> import numpy as np
>>> import pertpy as pt
>>> from anndata import AnnData
>>> effect_adata = AnnData(np.array([[-2.0, 1.0, 0.5]]))
>>> effect_adata.var_names = ["IL6", "CCR7", "other"]
>>> enr = pt.tl.Enrichment()

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Is it possible to use one of the existing datasets here instead of creating one? Makes the example shorter

>>> enr.signature_reversal(effect_adata, up_genes=["IL6"], down_genes=["CCR7"])
"""
if min_genes < 1:
raise ValueError("`min_genes` must be at least 1.")

matrix = _get_signature_matrix(adata, layer)
genes = _get_signature_genes(adata, gene_symbols_key)
query = _prepare_query_signature(query_signature, up_genes, down_genes)
aligned_query = query.reindex(genes).to_numpy(dtype=float)
present = np.isfinite(aligned_query)
n_present = int(present.sum())
if n_present < min_genes:
raise ValueError(f"Only {n_present} query genes are present in `adata`; at least {min_genes} are required.")

up_mask = aligned_query > 0
down_mask = aligned_query < 0
if (query > 0).any() and not up_mask.any():
raise ValueError("None of the query-up genes are present in `adata`.")
if (query < 0).any() and not down_mask.any():
raise ValueError("None of the query-down genes are present in `adata`.")

connectivity_scores = np.empty(adata.n_obs, dtype=float)
for idx in range(adata.n_obs):
values = np.asarray(to_dense(matrix[idx]), dtype=float).reshape(-1)
if not np.isfinite(values).all():
raise ValueError(
f"Perturbation signature {adata.obs_names[idx]!r} contains non-finite values; "
"all signatures must be finite and use the same gene universe."
)
connectivity_scores[idx] = _cmap_connectivity(values, up_mask, down_mask)

score_key = f"{key_added}_score"
connectivity_key = f"{key_added}_connectivity"
rank_key = f"{key_added}_rank"
reversal_scores = -connectivity_scores
adata.obs[score_key] = reversal_scores
adata.obs[connectivity_key] = connectivity_scores
adata.obs[rank_key] = (
pd.Series(reversal_scores, index=adata.obs_names)
.rank(ascending=False, method="min", na_option="keep")
.astype("Int64")
)
adata.uns[key_added] = {
"score_key": score_key,
"connectivity_key": connectivity_key,
"rank_key": rank_key,
"method": "cmap_wtcs",
"layer": layer,
"gene_symbols_key": gene_symbols_key,
"min_genes": min_genes,
"query_signature": query.to_dict(),
"up_genes": query.index[query > 0].tolist(),
"down_genes": query.index[query < 0].tolist(),
"matched_genes": genes[present].tolist(),
"n_query_genes": int(len(query)),
"n_matched_genes": n_present,
"n_matched_up_genes": int(up_mask.sum()),
"n_matched_down_genes": int(down_mask.sum()),
}

@deprecated_arg(
"pvals_adj_thresh",
Deprecation("1.0.6", "Use `padj_threshold`."),
Expand Down
Loading
Loading