Skip to content
Open
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
8 changes: 7 additions & 1 deletion openml/evaluations/__init__.py
Original file line number Diff line number Diff line change
@@ -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",
Expand Down
38 changes: 33 additions & 5 deletions openml/evaluations/functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand Down
46 changes: 46 additions & 0 deletions tests/test_evaluations/test_evaluation_functions.py
Original file line number Diff line number Diff line change
@@ -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


Expand Down Expand Up @@ -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()
Expand Down
Loading