diff --git a/openml/evaluations/__init__.py b/openml/evaluations/__init__.py index b56d0c2d5..29344b03a 100644 --- a/openml/evaluations/__init__.py +++ b/openml/evaluations/__init__.py @@ -1,10 +1,16 @@ # License: BSD 3-Clause from .evaluation import OpenMLEvaluation -from .functions import list_evaluation_measures, list_evaluations, list_evaluations_setups +from .functions import ( + list_estimation_procedures, + list_evaluation_measures, + list_evaluations, + list_evaluations_setups, +) __all__ = [ "OpenMLEvaluation", + "list_estimation_procedures", "list_evaluation_measures", "list_evaluations", "list_evaluations_setups", diff --git a/openml/evaluations/functions.py b/openml/evaluations/functions.py index f4e07c1b8..e9bfa3030 100644 --- a/openml/evaluations/functions.py +++ b/openml/evaluations/functions.py @@ -156,18 +156,46 @@ def list_evaluation_measures() -> list[str]: return openml._backend.evaluation_measure.list() -def list_estimation_procedures() -> list[str]: - """Return list of evaluation procedures available. +@overload +def list_estimation_procedures( + output_format: Literal["dataframe"], +) -> pd.DataFrame: ... + + +@overload +def list_estimation_procedures( + output_format: Literal["dict"] = ..., +) -> dict[int, dict[str, object]]: ... + + +def list_estimation_procedures( + output_format: Literal["dict", "dataframe"] = "dict", +) -> dict[int, dict[str, object]] | pd.DataFrame: + """Return the estimation procedures available on OpenML. The function performs an API call to retrieve the entire list of - evaluation procedures' names that are available. + evaluation procedures. Each procedure includes its ID, task type, + name, and type. + + Parameters + ---------- + output_format : {"dict", "dataframe"}, default="dict" + The format of the returned procedures. The dictionary format maps + procedure IDs to their remaining metadata. The DataFrame format has + one row per procedure, including an ``id`` column. Returns ------- - list + dict or pandas.DataFrame + The available estimation procedures in the requested format. """ result = openml._backend.estimation_procedure.list() - return [i.name for i in result] + records = [procedure._to_dict() for procedure in result] + + if output_format == "dataframe": + return pd.DataFrame.from_records(records) + + return {record.pop("id"): record for record in records} def list_evaluations_setups( diff --git a/tests/test_evaluations/test_evaluation_functions.py b/tests/test_evaluations/test_evaluation_functions.py index e15556d7b..dd4cea275 100644 --- a/tests/test_evaluations/test_evaluation_functions.py +++ b/tests/test_evaluations/test_evaluation_functions.py @@ -1,10 +1,14 @@ # License: BSD 3-Clause from __future__ import annotations +from unittest.mock import patch + import pytest import openml import openml.evaluations +from openml.estimation_procedures import OpenMLEstimationProcedure +from openml.tasks import TaskType from openml.testing import TestBase @@ -239,6 +243,48 @@ def test_list_evaluation_measures(self): assert isinstance(measures, list) is True assert all(isinstance(s, str) for s in measures) is True + def test_list_estimation_procedures_dict(self): + procedures = [ + OpenMLEstimationProcedure( + id=5, + task_type_id=TaskType.SUPERVISED_CLASSIFICATION, + name="10-fold Crossvalidation", + type="crossvalidation", + ) + ] + with patch.object(openml._backend.estimation_procedure, "list", return_value=procedures): + result = openml.evaluations.list_estimation_procedures() + + assert result == { + 5: { + "task_type_id": TaskType.SUPERVISED_CLASSIFICATION, + "name": "10-fold Crossvalidation", + "type": "crossvalidation", + } + } + + def test_list_estimation_procedures_dataframe(self): + procedures = [ + OpenMLEstimationProcedure( + id=5, + task_type_id=TaskType.SUPERVISED_CLASSIFICATION, + name="10-fold Crossvalidation", + type="crossvalidation", + ) + ] + with patch.object(openml._backend.estimation_procedure, "list", return_value=procedures): + result = openml.evaluations.list_estimation_procedures(output_format="dataframe") + + assert list(result.columns) == ["id", "task_type_id", "name", "type"] + assert result.to_dict("records") == [ + { + "id": 5, + "task_type_id": TaskType.SUPERVISED_CLASSIFICATION, + "name": "10-fold Crossvalidation", + "type": "crossvalidation", + } + ] + @pytest.mark.production_server() def test_list_evaluations_setups_filter_flow(self): self.use_production_server()