diff --git a/docs/_static/docstring_previews/dose_response.png b/docs/_static/docstring_previews/dose_response.png new file mode 100644 index 00000000..84a8e76b Binary files /dev/null and b/docs/_static/docstring_previews/dose_response.png differ diff --git a/docs/api/tools_index.md b/docs/api/tools_index.md index d421b960..aec4577c 100644 --- a/docs/api/tools_index.md +++ b/docs/api/tools_index.md @@ -553,3 +553,20 @@ similar = ds.nearest_perturbations(ds_adata, "IFNGR2", target_col="gene_target") ``` See [perturbation space tutorial](https://pertpy.readthedocs.io/en/latest/tutorials/notebooks/perturbation_space.html). + +### Dose-response curve fitting + +{meth}`~pertpy.tools.PseudobulkSpace.dose_response` returns an AnnData with one observation per perturbation and dose, holding the mean expression in `X` and the distance from control in `.obs`. +{meth}`~pertpy.tools.PseudobulkSpace.fit_dose_response` fits a four-parameter Hill curve per perturbation to the distance or any gene and stores EC50 and the other curve parameters in `.obs`. +For preprocessed `adata` with `perturbation` and `dose` columns, a `control` group and a PCA representation: + +```python +import pertpy as pt + +ps = pt.tl.PseudobulkSpace() +dose_adata = ps.dose_response(adata, embedding_key="X_pca") +ps.fit_dose_response(dose_adata) +ps.plot_dose_response(dose_adata) +``` + +Assay measurements such as viability can be fit the same way from an AnnData whose `.obs` holds perturbation, dose and response columns. diff --git a/src/pertpy/tools/_perturbation_space/_perturbation_space.py b/src/pertpy/tools/_perturbation_space/_perturbation_space.py index a5f66b0d..69abda0f 100644 --- a/src/pertpy/tools/_perturbation_space/_perturbation_space.py +++ b/src/pertpy/tools/_perturbation_space/_perturbation_space.py @@ -6,15 +6,21 @@ import numpy as np import pandas as pd +import scanpy as sc from anndata import AnnData +from scipy.optimize import curve_fit +from scipy.special import expit from scipy.stats import entropy +from pertpy._doc import _doc_params, doc_common_plot_args from pertpy._logger import logger from pertpy._types import cast_dense, cast_frame if TYPE_CHECKING: from collections.abc import Callable, Iterable, Sequence + from matplotlib.figure import Figure + from pertpy._types import RandomStateLike from pertpy.tools._distances._distances import Metric @@ -83,6 +89,71 @@ def _vector_distance(u: np.ndarray, v: np.ndarray, metric: str) -> float: raise ValueError(f"Unknown metric {metric!r}. Choose from 'euclidean', 'cosine', 'pearson'.") +def _four_parameter_logistic( + dose: np.ndarray, e0: float, emax: float, log_midpoint: float, hill_coefficient: float +) -> np.ndarray: + """Calculate predicted responses at the given doses from four-parameter Hill curve parameters.""" + log_dose = np.full_like(dose, -np.inf, dtype=float) + np.log(dose, out=log_dose, where=dose > 0) + fraction = expit(hill_coefficient * (log_dose - log_midpoint)) + return e0 + (emax - e0) * fraction + + +def _fit_hill(doses: np.ndarray, responses: np.ndarray) -> dict[str, float | bool]: + """Fit a four-parameter Hill curve to one perturbation's responses.""" + unique_doses, dose_index = np.unique(doses, return_inverse=True) + if unique_doses.size < 4: + raise ValueError("at least four distinct dose values are required.") + response_offset = float(np.min(responses)) + response_scale = float(np.ptp(responses)) + if response_scale == 0: + raise ValueError("the response is constant.") + + responses = (responses - response_offset) / response_scale + mean_response = np.bincount(dose_index, weights=responses) / np.bincount(dose_index) + e0_guess, emax_guess = float(mean_response[0]), float(mean_response[-1]) + positive = unique_doses > 0 + halfway_distance = np.abs(mean_response[positive] - (e0_guess + emax_guess) / 2) + midpoint_guess = float(unique_doses[positive][np.argmin(halfway_distance)]) + + parameters, covariance = curve_fit( + _four_parameter_logistic, + doses, + responses, + p0=(e0_guess, emax_guess, np.log(midpoint_guess), 1.0), + bounds=((-np.inf, -np.inf, -np.inf, np.finfo(float).eps), np.inf), + absolute_sigma=True, + maxfev=20_000, + ) + e0, emax, log_midpoint, hill_coefficient = parameters + midpoint = float(np.exp(log_midpoint)) + residual_sum_squares = float(np.sum((responses - _four_parameter_logistic(doses, *parameters)) ** 2)) + total_sum_squares = float(np.sum((responses - responses.mean()) ** 2)) + positive_doses = doses[doses > 0] + + # Check conditioning before scaling covariance by residual variance, which may be zero. + degrees_of_freedom = len(responses) - len(parameters) + if ( + degrees_of_freedom <= 0 + or not np.isfinite(covariance).all() + or np.linalg.matrix_rank(covariance) < len(parameters) + ): + midpoint_standard_error = np.nan + else: + residual_variance = residual_sum_squares / degrees_of_freedom + midpoint_standard_error = midpoint * np.sqrt(float(covariance[2, 2]) * residual_variance) + + return { + "e0": float(e0 * response_scale + response_offset), + "emax": float(emax * response_scale + response_offset), + "slope": float(hill_coefficient), + "ec50": midpoint, + "ec50_se": float(midpoint_standard_error), + "r_squared": 1 - residual_sum_squares / total_sum_squares, + "ec50_in_range": bool(positive_doses.min() < midpoint < positive_doses.max()), + } + + def _subtract_control_mean( matrix: np.ndarray, control_mask: np.ndarray, @@ -590,7 +661,7 @@ def dose_response( layer_key: str | None = None, embedding_key: str | None = None, **kwargs, - ) -> pd.DataFrame: + ) -> AnnData: """Quantify the effect size of each perturbation as a function of dose. For every (perturbation, dose) group the statistical distance to ``reference_key`` is computed in the chosen representation using :class:`~pertpy.tools.Distance`. @@ -602,18 +673,18 @@ def dose_response( dose_col: `.obs` column with the (numeric) dose. reference_key: Control perturbation all doses are compared against. metric: Distance metric passed to :class:`~pertpy.tools.Distance`. - layer_key: Layer to compute distances from. + layer_key: Layer to compute distances and mean expression from. embedding_key: `.obsm` embedding to compute distances from. kwargs: Passed to :meth:`~pertpy.tools.Distance.onesided_distances`. Returns: - Tidy DataFrame with ``perturbation``, ``dose`` and ``distance`` columns, sorted by perturbation then dose. + AnnData with one observation per non-reference (perturbation, dose) group, holding the mean expression in ``X`` and the ``distance`` in `.obs`. Examples: >>> import pertpy as pt >>> adata = pt.ds.srivatsan_2020_sciplex2() >>> ps = pt.tl.PseudobulkSpace() - >>> curves = ps.dose_response(adata, dose_col="dose_value", embedding_key="X_pca") + >>> dose_adata = ps.dose_response(adata, dose_col="dose_value", embedding_key="X_pca") """ for col in (target_col, dose_col): if col not in adata.obs: @@ -635,18 +706,166 @@ def dose_response( if isinstance(dists, tuple): dists = dists[0] - records = [] - for label, value in dists.items(): - if label == reference_key: - continue - perturbation, _, dose = str(label).partition(sep) - records.append({"perturbation": perturbation, "dose": dose, "distance": float(value)}) - result = pd.DataFrame.from_records(records) + treated = grouped[~is_control].copy() + treated.obs["_dose_group"] = treated.obs["_dose_group"].cat.remove_unused_categories() + dose_adata = sc.get.aggregate(treated, by="_dose_group", func="mean", layer=layer_key) + dose_adata.X = dose_adata.layers.pop("mean") + _carry_constant_obs(dose_adata, cast_frame(treated.obs), "_dose_group") + dose_obs = cast_frame(dose_adata.obs) + dose_adata.obs["distance"] = dists.reindex(dose_obs["_dose_group"].astype(str)).to_numpy(dtype=float) with warnings.catch_warnings(): warnings.simplefilter("ignore") with contextlib.suppress(ValueError, TypeError): - result["dose"] = pd.to_numeric(result["dose"]) - return result.sort_values(["perturbation", "dose"]).reset_index(drop=True) + dose_adata.obs[dose_col] = pd.to_numeric(dose_obs[dose_col].astype(str)) + dose_adata.obs = cast_frame(dose_adata.obs).drop(columns="_dose_group") + dose_adata.obs_names = dose_adata.obs_names.str.replace(sep, "_") + order = cast_frame(dose_adata.obs).sort_values([target_col, dose_col]).index + return dose_adata[order].copy() + + def fit_dose_response( + self, + adata: AnnData, + response: str = "distance", + *, + target_col: str = "perturbation", + dose_col: str = "dose", + key_added: str | None = None, + ) -> None: + """Fit a four-parameter Hill curve for each perturbation. + + Perturbations whose curve cannot be fit get a warning and NaN parameters. + + Args: + adata: AnnData with one observation per dose or replicate, such as the output of :meth:`dose_response`. + response: `.obs` column or gene in ``var_names`` to fit. + target_col: `.obs` column identifying the perturbation. + dose_col: `.obs` column containing non-negative numeric doses. + key_added: Prefix of the `.obs` columns the results are written to. Defaults to ``response``. + + Returns: + Adds the fitted response and the per-perturbation ``e0``, ``emax``, ``slope``, ``ec50``, ``ec50_se``, ``r_squared`` and ``ec50_in_range`` to `.obs`, prefixed with ``key_added``. + + Examples: + >>> import pertpy as pt + >>> adata = pt.ds.srivatsan_2020_sciplex2() + >>> ps = pt.tl.PseudobulkSpace() + >>> dose_adata = ps.dose_response(adata, dose_col="dose_value", embedding_key="X_pca") + >>> ps.fit_dose_response(dose_adata, dose_col="dose_value") + >>> ps.fit_dose_response(dose_adata, "CDKN1A", dose_col="dose_value") + """ + data = sc.get.obs_df(adata, keys=[target_col, dose_col, response]) + doses = data[dose_col].to_numpy(dtype=float) + responses = data[response].to_numpy(dtype=float) + if (doses < 0).any(): + raise ValueError("Dose values must be non-negative.") + + labels = data[target_col].to_numpy() + fitted = np.full(len(data), np.nan) + records: dict[object, dict[str, float | bool]] = {} + for perturbation in pd.unique(labels): + mask = labels == perturbation + try: + fit = _fit_hill(doses[mask], responses[mask]) + except (ValueError, RuntimeError) as e: + warnings.warn( + f"Cannot fit a Hill curve for perturbation {perturbation!r}: {e}", UserWarning, stacklevel=2 + ) + fit = dict.fromkeys(("e0", "emax", "slope", "ec50", "ec50_se", "r_squared"), np.nan) + fit["ec50_in_range"] = False + else: + fitted[mask] = _four_parameter_logistic( + doses[mask], fit["e0"], fit["emax"], np.log(fit["ec50"]), fit["slope"] + ) + if np.isnan(fit["ec50_se"]): + warnings.warn( + f"Cannot estimate the EC50 standard error for perturbation {perturbation!r}. " + "Inspect the dose range and fitted curve before interpreting the estimate.", + UserWarning, + stacklevel=2, + ) + records[perturbation] = fit + + prefix = response if key_added is None else key_added + fits = pd.DataFrame.from_dict(records, orient="index") + adata.obs[f"{prefix}_fitted"] = fitted + for col in fits.columns: + adata.obs[f"{prefix}_{col}"] = fits[col].reindex(labels).to_numpy() + + @_doc_params(common_plot_args=doc_common_plot_args) + def plot_dose_response( # pragma: no cover # noqa: D417 + self, + adata: AnnData, + response: str = "distance", + *, + target_col: str = "perturbation", + dose_col: str = "dose", + key_added: str | None = None, + perturbations: Sequence[str] | None = None, + ncols: int = 4, + return_fig: bool = False, + ) -> Figure | None: + """Plot the measured responses and fitted Hill curves of each perturbation. + + Args: + adata: AnnData with one observation per dose or replicate, such as the output of :meth:`dose_response`. + response: `.obs` column or gene in ``var_names`` to plot. + target_col: `.obs` column identifying the perturbation. + dose_col: `.obs` column containing the doses. + key_added: Prefix passed to :meth:`fit_dose_response`. Defaults to ``response``. + perturbations: Perturbations to plot. Defaults to all. + ncols: Number of panels per row. + {common_plot_args} + + Returns: + If `return_fig` is `True`, returns the figure, otherwise `None`. + + Examples: + >>> import pertpy as pt + >>> adata = pt.ds.srivatsan_2020_sciplex2() + >>> ps = pt.tl.PseudobulkSpace() + >>> dose_adata = ps.dose_response(adata, dose_col="dose_value", embedding_key="X_pca") + >>> ps.fit_dose_response(dose_adata, dose_col="dose_value") + >>> ps.plot_dose_response(dose_adata, dose_col="dose_value") + + Preview: + .. image:: /_static/docstring_previews/dose_response.png + """ + import matplotlib.pyplot as plt + + prefix = response if key_added is None else key_added + data = sc.get.obs_df(adata, keys=[target_col, dose_col, response]) + obs = cast_frame(adata.obs) + labels = list(pd.unique(data[target_col]) if perturbations is None else perturbations) + ncols = min(ncols, len(labels)) + nrows = -(-len(labels) // ncols) + fig, axes = plt.subplots(nrows, ncols, figsize=(3.5 * ncols, 3 * nrows), squeeze=False, layout="constrained") + for ax, label in zip(axes.flat, labels, strict=False): + mask = (data[target_col] == label).to_numpy() + doses = data.loc[mask, dose_col].to_numpy(dtype=float) + ax.scatter(doses, data.loc[mask, response], zorder=3) + if f"{prefix}_ec50" in obs and np.isfinite((fit := obs.loc[mask].iloc[0])[f"{prefix}_ec50"]): + grid = np.geomspace(doses[doses > 0].min(), doses.max(), 200) + curve = _four_parameter_logistic( + grid, + fit[f"{prefix}_e0"], + fit[f"{prefix}_emax"], + np.log(fit[f"{prefix}_ec50"]), + fit[f"{prefix}_slope"], + ) + ax.plot(grid, curve, color="tab:orange") + if fit[f"{prefix}_ec50_in_range"]: + ax.axvline(fit[f"{prefix}_ec50"], color="0.5", linestyle=":") + lowest = doses[doses > 0].min() + ax.set_xscale("symlog", linthresh=lowest) + ax.set_xlim(0 if (doses == 0).any() else lowest / 2, doses.max() * 2) + ax.set(title=str(label), xlabel=dose_col, ylabel=response) + for ax in axes.flat[len(labels) :]: + ax.set_visible(False) + + if return_fig: + return fig + plt.show() + return None def plot_similarity( # pragma: no cover self, diff --git a/tests/tools/_perturbation_space/test_perturbation_space_extras.py b/tests/tools/_perturbation_space/test_perturbation_space_extras.py index 3d1501e8..db084706 100644 --- a/tests/tools/_perturbation_space/test_perturbation_space_extras.py +++ b/tests/tools/_perturbation_space/test_perturbation_space_extras.py @@ -5,6 +5,7 @@ from anndata import AnnData import pertpy as pt +from pertpy._types import cast_frame @pytest.fixture @@ -66,19 +67,98 @@ def test_evaluate_combinations(rng): np.testing.assert_allclose(result.loc["A+B", "distance"], 0.0, atol=1e-6) -def test_dose_response(rng): +@pytest.mark.parametrize("categorical_doses", [False, True]) +def test_dose_response(rng, categorical_doses): groups, doses = [], [] for pert in ["control", "drug"]: - for dose in [0.0] if pert == "control" else [1.0, 10.0, 100.0]: + for dose in [0.0] if pert == "control" else [0.1, 1.0, 3.0, 10.0, 30.0, 100.0]: groups += [pert] * 15 doses += [dose] * 15 groups = np.array(groups) doses = np.array(doses, dtype=float) - X = rng.normal(0, 0.3, (len(groups), 8)) - X[groups == "drug"] += (doses[groups == "drug"] / 10.0)[:, None] + X = np.tile(rng.normal(0, 0.3, (15, 8)), (len(groups) // 15, 1)) + drug_doses = doses[groups == "drug"] + X[groups == "drug"] += (5 * drug_doses**1.2 / (10**1.2 + drug_doses**1.2))[:, None] adata = AnnData(X, obs=pd.DataFrame({"perturbation": groups, "dose": doses})) + if categorical_doses: + adata.obs["dose"] = pd.Categorical(adata.obs["dose"].astype(str)) sc.pp.pca(adata, n_comps=5) - curves = pt.tl.PseudobulkSpace().dose_response(adata, dose_col="dose", metric="euclidean", embedding_key="X_pca") - drug = curves[curves["perturbation"] == "drug"].sort_values("dose") - assert drug["distance"].is_monotonic_increasing + ps = pt.tl.PseudobulkSpace() + dose_adata = ps.dose_response(adata, dose_col="dose", metric="euclidean", embedding_key="X_pca") + assert dose_adata.shape == (6, 8) + assert dose_adata.obs["dose"].tolist() == [0.1, 1.0, 3.0, 10.0, 30.0, 100.0] + assert dose_adata.obs["distance"].is_monotonic_increasing + + ps.fit_dose_response(dose_adata) + ps.fit_dose_response(dose_adata, dose_adata.var_names[0]) + assert dose_adata.obs["distance_ec50"].iloc[0] == pytest.approx(10, rel=1e-4) + assert dose_adata.obs["distance_ec50_in_range"].all() + assert dose_adata.obs[f"{dose_adata.var_names[0]}_ec50"].iloc[0] == pytest.approx(10, rel=1e-4) + + +def _assay(data: pd.DataFrame) -> AnnData: + return AnnData(obs=data.reset_index(drop=True).rename(index=str)) + + +def _fits(adata: AnnData) -> pd.DataFrame: + return cast_frame(adata.obs).drop_duplicates("perturbation").set_index("perturbation") + + +@pytest.mark.parametrize(("e0", "emax"), [(0.1, 1.8), (1.0, 0.05)]) +def test_fit_dose_response(e0, emax): + doses = np.array([0.0, 0.1, 0.3, 1.0, 3.0, 10.0, 30.0, 100.0]) + responses = e0 + (emax - e0) * doses**1.4 / (3.0**1.4 + doses**1.4) + adata = _assay(pd.DataFrame({"perturbation": "drug", "dose": doses, "distance": responses})) + pt.tl.PseudobulkSpace().fit_dose_response(adata) + fit = _fits(adata).loc["drug"] + + assert fit["distance_e0"] == pytest.approx(e0) + assert fit["distance_emax"] == pytest.approx(emax) + assert fit["distance_slope"] == pytest.approx(1.4) + assert fit["distance_ec50"] == pytest.approx(3.0) + np.testing.assert_allclose(adata.obs["distance_fitted"], responses, rtol=1e-6) + + +def test_fit_dose_response_standard_error(): + doses = np.array([0.0, 0.1, 0.3, 1.0, 3.0, 10.0, 30.0, 100.0, 300.0]) + responses = doses / (10 + doses) + np.random.default_rng(2026).normal(0, 0.025, len(doses)) + complete = pd.DataFrame({"perturbation": "complete", "dose": doses, "distance": responses}) + limited = complete.loc[complete["dose"] <= 3].assign(perturbation="limited") + adata = _assay(pd.concat([complete, limited])) + + with pytest.warns(UserWarning, match="Cannot estimate.*'limited'"): + pt.tl.PseudobulkSpace().fit_dose_response(adata) + fits = _fits(adata) + fit = fits.loc["complete"] + + midpoint, slope = fit["distance_ec50"], fit["distance_slope"] + fraction = doses**slope / (midpoint**slope + doses**slope) + log_ratio = np.zeros_like(doses) + np.log(doses / midpoint, out=log_ratio, where=doses > 0) + sensitivity = (fit["distance_emax"] - fit["distance_e0"]) * fraction * (1 - fraction) + jacobian = np.column_stack((1 - fraction, fraction, -slope * sensitivity / midpoint, sensitivity * log_ratio)) + residuals = responses - (fit["distance_e0"] + (fit["distance_emax"] - fit["distance_e0"]) * fraction) + variance = np.sum(residuals**2) / (len(doses) - 4) + assert fit["distance_ec50_se"] == pytest.approx( + np.sqrt(np.linalg.inv(jacobian.T @ jacobian)[2, 2] * variance), rel=1e-4 + ) + assert np.isnan(fits.loc["limited", "distance_ec50_se"]) + + +def test_fit_dose_response_unfittable(): + doses = np.array([0.0, 0.1, 1.0, 3.0, 10.0, 30.0, 100.0]) + good = pd.DataFrame({"perturbation": "good", "dose": doses, "distance": doses / (10 + doses)}) + adata = _assay(pd.concat([good, good.assign(perturbation="bad", distance=1.0)])) + + with pytest.warns(UserWarning, match="'bad'.*constant"): + pt.tl.PseudobulkSpace().fit_dose_response(adata) + fits = _fits(adata) + assert fits.loc["good", "distance_ec50"] == pytest.approx(10) + assert np.isnan(fits.loc["bad", "distance_ec50"]) + + +def test_fit_dose_response_negative_dose(): + adata = _assay(pd.DataFrame({"perturbation": "drug", "dose": [-1, 1, 2, 3], "distance": [0, 1, 2, 3]})) + with pytest.raises(ValueError, match="non-negative"): + pt.tl.PseudobulkSpace().fit_dose_response(adata)