"""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