Source code for scplotkit.ridgeline

"""Ridgeline (joy) plots: stacked, overlapping KDEs of a continuous value across groups."""

from __future__ import annotations

from collections.abc import Sequence
from pathlib import Path

import matplotlib.pyplot as plt
import numpy as np
import pandas as pd
from matplotlib.patches import Patch
from scipy.stats import gaussian_kde

from ._config import PlotConfig
from ._palettes import resolve_palette
from ._style import save_figure
from ._utils import require_columns

__all__ = ["ridgeline_plot"]


[docs] def ridgeline_plot( adata, obs_key: str, group_by: str, condition_column: str | None = None, celltype_column: str | None = None, celltype_of_interest: str | Sequence | None = None, scale: bool = False, palette: dict | None = None, bw_adjust: float = 1.0, overlap: float = 0.6, figsize: tuple | None = None, alpha: float = 0.72, show_median: bool = True, title: str | None = None, save_name: str | None = None, config: PlotConfig | None = None, output_dir: str | Path | None = "figures", ): """Ridgeline plot of a continuous ``.obs`` column, one ridge per ``group_by`` value. Parameters ---------- obs_key: Continuous ``.obs`` column to plot (x-axis). group_by: Categorical ``.obs`` column defining the rows (one ridge per value). condition_column: If given, overlays one KDE per condition value on each ridge, colored by condition (looked up via ``palette``/config instead of ``group_by``). celltype_column, celltype_of_interest: Optional subsetting before computing distributions. scale: Z-score ``obs_key`` across all cells before plotting. bw_adjust: Bandwidth multiplier for the KDE (like R's ``adjust=``). overlap: Fraction of row height that ridges overlap (``0`` = no overlap). show_median: Draw a vertical tick at the median of each distribution. """ config = PlotConfig.load(config) require_columns(adata, obs_key, group_by, condition_column, celltype_column) general = config.general if celltype_column and celltype_of_interest is not None: values = [celltype_of_interest] if isinstance(celltype_of_interest, str) else celltype_of_interest adata = adata[adata.obs[celltype_column].isin(values)].copy() cols = [obs_key, group_by] + ([condition_column] if condition_column else []) df = adata.obs[cols].copy() df[obs_key] = pd.to_numeric(df[obs_key], errors="coerce") df = df.dropna(subset=[obs_key]) if scale: mu, sigma = df[obs_key].mean(), df[obs_key].std() df[obs_key] = (df[obs_key] - mu) / sigma groups = df[group_by].cat.categories.tolist() if hasattr(df[group_by], "cat") else sorted( df[group_by].unique().tolist() ) n_groups = len(groups) if condition_column: conditions = ( df[condition_column].cat.categories.tolist() if hasattr(df[condition_column], "cat") else sorted(df[condition_column].unique().tolist()) ) colors = resolve_palette(config, condition_column, conditions, palette) else: conditions = None colors = resolve_palette(config, group_by, groups, palette) x_lo, x_hi = df[obs_key].quantile(0.001), df[obs_key].quantile(0.999) x_vals = np.linspace(x_lo, x_hi, 600) kde_cache = {} max_density = 0.0 for g in groups: gdf = df[df[group_by] == g] iter_keys = [(g, c) for c in conditions] if conditions else [g] for key in iter_keys: vals = gdf[gdf[condition_column] == key[1]][obs_key].values if conditions else gdf[obs_key].values if len(vals) < 3: continue kde = gaussian_kde(vals) kde.set_bandwidth(kde.factor * bw_adjust) density = kde(x_vals) max_density = max(max_density, density.max()) kde_cache[key] = (density, np.median(vals)) if max_density == 0: raise ValueError("No groups had enough data to compute a KDE.") row_height = max_density * (1.0 + overlap) if figsize is None: figsize = (8, max(3, n_groups * 0.75 + 1.5)) fig, ax = plt.subplots(figsize=figsize) for i, g in enumerate(groups): y_base = i * row_height iter_keys = [(g, c) for c in conditions] if conditions else [g] for key in iter_keys: if key not in kde_cache: continue density, median_val = kde_cache[key] y = density + y_base color = colors[key[1]] if conditions else colors[g] ax.fill_between(x_vals, y_base, y, alpha=alpha, color=color, lw=0) ax.plot(x_vals, y, color=color, lw=1.4, alpha=min(alpha + 0.2, 1.0)) if show_median: med_density = kde_cache[key][0][np.argmin(np.abs(x_vals - median_val))] ax.plot( [median_val, median_val], [y_base, y_base + med_density], color="white", lw=1.5, ls="--", alpha=0.85, zorder=5, ) ax.axhline(y_base, color="white", lw=1.0, zorder=4) ax.set_yticks([i * row_height for i in range(n_groups)]) ax.set_yticklabels(groups, fontsize=general["legend_fontsize"]) ax.yaxis.set_tick_params(length=0) ax.set_ylim(-row_height * 0.05, n_groups * row_height) x_label = obs_key.replace("_", " ") if scale: x_label += " (z-scored)" ax.set_xlabel(x_label, fontsize=general["legend_fontsize"]) ax.set_xlim(x_lo, x_hi) for sp in ("top", "right", "left"): ax.spines[sp].set_visible(False) ax.set_title(title or obs_key.replace("_", " "), fontsize=general["title_fontsize"]) ax.grid(False) if conditions: handles = [Patch(facecolor=colors[c], alpha=alpha, label=str(c)) for c in conditions] ax.legend( handles=handles, bbox_to_anchor=(1.01, 1), loc="upper left", frameon=False, fontsize=general["legend_fontsize"], title=condition_column.replace("_", " ").title(), ) plt.tight_layout() if save_name is not None: save_figure(fig, output_dir, "ridgeline", f"{save_name}.png", config) return fig