Source code for scplotkit.enrichment

"""Over-representation analysis (ORA) dot plots, e.g. from Enrichr/gseapy results."""

from __future__ import annotations

from pathlib import Path

import matplotlib.pyplot as plt
import numpy as np
import pandas as pd

from ._config import PlotConfig
from ._style import save_figure

__all__ = ["ora_dotplot"]

_X_METRICS = {
    "gene_ratio": ("gene_ratio", "Gene Ratio"),
    "odds_ratio": ("Odds Ratio", "Odds Ratio"),
    "combined_score": ("Combined Score", "Combined Score"),
}


[docs] def ora_dotplot( csv_path: str | Path, n_top: int = 20, gene_set_filter: str | list[str] | None = None, padj_threshold: float = 0.05, x_metric: str = "gene_ratio", cmap: str = "magma", figsize: tuple | None = None, title: str | None = None, save_name: str | None = None, config: PlotConfig | None = None, output_dir: str | Path | None = "figures", ): """Dot plot overview of Over-Representation Analysis (ORA) results. Expects a CSV (e.g. exported from Enrichr/gseapy) with columns ``Gene_set``, ``Term``, ``Overlap`` (``"n/total"`` strings), ``Adjusted P-value``, ``Odds Ratio``, ``Combined Score``. Dot size scales with the number of overlapping genes; dot color encodes ``-log10(adjusted p-value)``. Parameters ---------- n_top: Number of top terms to display, ranked by ``Combined Score``. gene_set_filter: Restrict to specific ``Gene_set`` value(s). padj_threshold: Only show terms with adjusted p-value below this cutoff. x_metric: ``'gene_ratio'``, ``'odds_ratio'``, or ``'combined_score'``. """ config = PlotConfig.load(config) general = config.general df = pd.read_csv(csv_path) df[["n_overlap", "pathway_size"]] = df["Overlap"].str.split("/", expand=True).astype(int) df["gene_ratio"] = df["n_overlap"] / df["pathway_size"] df["neg_log10_padj"] = -np.log10(df["Adjusted P-value"].clip(lower=1e-300)) if gene_set_filter is not None: gene_set_filter = [gene_set_filter] if isinstance(gene_set_filter, str) else gene_set_filter df = df[df["Gene_set"].isin(gene_set_filter)] df = df[df["Adjusted P-value"] < padj_threshold] if len(df) == 0: print("No significant terms to plot after filtering.") return None df = df.sort_values("Combined Score", ascending=False).head(n_top).sort_values( "Combined Score", ascending=True ).reset_index(drop=True) x_col, x_label = _X_METRICS.get(x_metric, _X_METRICS["gene_ratio"]) x = df[x_col] size_min, size_max = 30, 280 n_max = df["n_overlap"].max() sizes = df["n_overlap"] / n_max * (size_max - size_min) + size_min if figsize is None: figsize = (9, max(4, len(df) * 0.38 + 1.5)) fig, ax = plt.subplots(figsize=figsize) scatter = ax.scatter( x, np.arange(len(df)), c=df["neg_log10_padj"], s=sizes, cmap=cmap, vmin=df["neg_log10_padj"].min(), vmax=df["neg_log10_padj"].max(), zorder=3, edgecolors="none", ) cbar = fig.colorbar(scatter, ax=ax, shrink=0.35, pad=0.02, aspect=18) cbar.set_label("-log₁₀(adj. p-value)", fontsize=general["legend_fontsize"]) cbar.ax.tick_params(labelsize=general["legend_fontsize"] - 1) cbar.outline.set_visible(False) legend_vals = sorted({1, n_max // 2, n_max}) size_handles = [ plt.scatter([], [], s=v / n_max * (size_max - size_min) + size_min, color="#888888", edgecolors="none", label=str(v)) for v in legend_vals ] size_legend = ax.legend( handles=size_handles, title="Overlap\ngenes", bbox_to_anchor=(1.22, 0.0), loc="lower left", frameon=False, fontsize=general["legend_fontsize"], labelspacing=1.1, ) ax.add_artist(size_legend) ax.set_axisbelow(True) ax.xaxis.grid(True, color="#eeeeee", linewidth=0.7) ax.yaxis.grid(False) ax.set_yticks(np.arange(len(df))) ax.set_yticklabels(df["Term"], fontsize=general["legend_fontsize"]) ax.set_xlabel(x_label, fontsize=general["legend_fontsize"]) ax.set_ylim(-0.6, len(df) - 0.4) for spine in ("top", "right", "left"): ax.spines[spine].set_visible(False) if title: ax.set_title(title, fontsize=general["title_fontsize"]) elif gene_set_filter: label = gene_set_filter[0] if isinstance(gene_set_filter, list) else gene_set_filter ax.set_title(label.replace("_", " "), fontsize=general["title_fontsize"]) plt.tight_layout() if save_name is not None: save_figure(fig, output_dir, "ora", f"{save_name}.png", config) return fig