Skip to content
Binary file added docs/_static/docstring_previews/dose_response.png
Loading
Sorry, something went wrong. Reload?
Sorry, we cannot display this file.
Sorry, this file is invalid so it cannot be displayed.
17 changes: 17 additions & 0 deletions docs/api/tools_index.md
Original file line number Diff line number Diff line change
Expand Up @@ -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).
Comment thread
Zethson marked this conversation as resolved.

### 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.
245 changes: 232 additions & 13 deletions src/pertpy/tools/_perturbation_space/_perturbation_space.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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`.
Expand All @@ -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:
Expand All @@ -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:
Comment thread
Zethson marked this conversation as resolved.
>>> 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,
Expand Down
Loading
Loading