From 7b108ab001c851c749458180a0bafdc98acd284b Mon Sep 17 00:00:00 2001 From: Sizerta Date: Sun, 20 Sep 2026 03:10:24 +0400 Subject: [PATCH] fix(augur): use regression-appropriate scorers for regressor estimators --- src/pertpy/tools/_augur.py | 34 ++++++++++++++++------------------ tests/tools/test_augur.py | 19 +++++++++++++++++++ 2 files changed, 35 insertions(+), 18 deletions(-) diff --git a/src/pertpy/tools/_augur.py b/src/pertpy/tools/_augur.py index eb7cbd5d..7b5a63d8 100644 --- a/src/pertpy/tools/_augur.py +++ b/src/pertpy/tools/_augur.py @@ -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"), @@ -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, diff --git a/tests/tools/test_augur.py b/tests/tools/test_augur.py index 46de2b07..6d805083 100644 --- a/tests/tools/test_augur.py +++ b/tests/tools/test_augur.py @@ -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 @@ -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)