Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
34 changes: 16 additions & 18 deletions src/pertpy/tools/_augur.py
Original file line number Diff line number Diff line change
Expand Up @@ -471,6 +471,14 @@ def set_scorer(
>>> ag_rfc = pt.tl.Augur("random_forest_classifier")
>>> scorer = ag_rfc.set_scorer(True, 0)
"""
if is_regressor(self.estimator):
return {
"augur_score": make_scorer(self.ccc_score),
"r2": make_scorer(r2_score),
"ccc": make_scorer(self.ccc_score),
"neg_mean_squared_error": make_scorer(root_mean_squared_error),
"explained_variance": make_scorer(explained_variance_score),
}
if multiclass:
return {
"augur_score": make_scorer(roc_auc_score, multi_class="ovo", response_method="predict_proba"),
Expand All @@ -480,24 +488,14 @@ def set_scorer(
"f1": make_scorer(f1_score, average="macro"),
"recall": make_scorer(recall_score, average="macro"),
}
return (
{
"augur_score": make_scorer(roc_auc_score, response_method="predict_proba"),
"auc": make_scorer(roc_auc_score, response_method="predict_proba"),
"accuracy": make_scorer(accuracy_score),
"precision": make_scorer(precision_score, average="binary", zero_division=zero_division),
"f1": make_scorer(f1_score, average="binary"),
"recall": make_scorer(recall_score, average="binary"),
}
if isinstance(self.estimator, RandomForestClassifier | LogisticRegression)
else {
"augur_score": make_scorer(self.ccc_score),
"r2": make_scorer(r2_score),
"ccc": make_scorer(self.ccc_score),
"neg_mean_squared_error": make_scorer(root_mean_squared_error),
"explained_variance": make_scorer(explained_variance_score),
}
)
return {
"augur_score": make_scorer(roc_auc_score, response_method="predict_proba"),
"auc": make_scorer(roc_auc_score, response_method="predict_proba"),
"accuracy": make_scorer(accuracy_score),
"precision": make_scorer(precision_score, average="binary", zero_division=zero_division),
"f1": make_scorer(f1_score, average="binary"),
"recall": make_scorer(recall_score, average="binary"),
}

def run_cross_validation(
self,
Expand Down
19 changes: 19 additions & 0 deletions tests/tools/test_augur.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,10 @@
from pathlib import Path

import numpy as np
import pandas as pd
import pytest
import scanpy as sc
from anndata import AnnData

import pertpy as pt

Expand Down Expand Up @@ -97,6 +99,23 @@ def test_regressor(adata):
assert any([isclose(cv["mean_ccc"], ccc, abs_tol=10**-5), isclose(cv["mean_r2"], r2, abs_tol=10**-5)])


def test_regressor_scorer_with_more_than_two_labels():
"""Regressors get regression scorers even when the target has more than two values (#655)."""
scorer = ag_rfr.set_scorer(multiclass=True, zero_division=0)
assert set(scorer) == {"augur_score", "r2", "ccc", "neg_mean_squared_error", "explained_variance"}


def test_regressor_cross_validation_with_four_timepoints():
"""A regressor on four numeric timepoints must not request class probabilities (#655)."""
rng = np.random.default_rng(0)
y = np.repeat([0.0, 1.0, 2.0, 3.0], 15)
x = rng.poisson(2 + y[:, None] * (np.arange(30) < 5)).astype(float)
adata = AnnData(x, obs=pd.DataFrame({"y_": y}, index=[f"c{i}" for i in range(len(y))]))
cv = ag_rfr.run_cross_validation(adata, subsample_idx=0, folds=3, random_state=42, zero_division=0)
assert np.isfinite(cv["mean_ccc"])
assert "mean_auc" not in cv


def test_subsample(adata):
"""Test default, permute and velocity subsampling process."""
adata = ag_rfc.load(adata)
Expand Down
Loading