diff --git a/docs/api.rst b/docs/api.rst index aadd2718..18f8bc72 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -44,6 +44,7 @@ Preferential Optimization optuna_dashboard.preferential.create_study optuna_dashboard.preferential.load_study optuna_dashboard.preferential.PreferentialStudy + optuna_dashboard.register_preference_feedback_component Streamlit ----------------- diff --git a/optuna_dashboard/__init__.py b/optuna_dashboard/__init__.py index 5bb3f301..ea2a8dbe 100644 --- a/optuna_dashboard/__init__.py +++ b/optuna_dashboard/__init__.py @@ -14,6 +14,7 @@ from ._form_widget import TextInputWidget # noqa from ._named_objectives import set_objective_names # noqa from ._note import get_note # noqa from ._note import save_note # noqa +from ._preference_setting import register_preference_feedback_component # noqa __version__ = "0.13.0b1" diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 731de524..abd717d0 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -28,6 +28,7 @@ from ._cached_extra_study_property import get_cached_extra_study_property from ._custom_plot_data import get_plotly_graph_objects from ._importance import get_param_importance_from_trials_cache from ._pareto_front import get_pareto_front_trials +from ._preference_setting import _register_preference_feedback_component from ._preferential_history import NewHistory from ._preferential_history import PreferenceHistoryNotFound from ._preferential_history import remove_history @@ -315,6 +316,28 @@ def create_app( response.status = 204 return {} + @app.put("/api/studies//preference_feedback_component_type") + @json_api_view + def put_preference_feedback_component_type(study_id: int) -> dict[str, Any]: + try: + component_type = request.json.get("type", "") + artifact_key = request.json.get("artifact_key", None) + except ValueError: + response.status = 400 + return {"reason": "invalid request."} + if component_type not in ["note", "artifact"]: + response.status = 400 + return {"reason": "component_type must be either 'Note' or 'Artifact'."} + + _register_preference_feedback_component( + study_id=study_id, + storage=storage, + component_type=component_type, + artifact_key=artifact_key, + ) + response.status = 204 + return {} + @app.delete("/api/studies//preference/") @json_api_view def remove_preference(study_id: int, history_id: str) -> dict[str, Any]: diff --git a/optuna_dashboard/_preference_setting.py b/optuna_dashboard/_preference_setting.py new file mode 100644 index 00000000..907fc2be --- /dev/null +++ b/optuna_dashboard/_preference_setting.py @@ -0,0 +1,67 @@ +from __future__ import annotations + +from typing import Any +from typing import TYPE_CHECKING + +from optuna.storages import BaseStorage + +from .preferential._study import PreferentialStudy + + +if TYPE_CHECKING: + from typing import Literal + + OUTPUT_COMPONENT_TYPE = Literal["note", "artifact"] + +_SYSTEM_ATTR_FEEDBACK_COMPONENT = "preference:component" + + +def _register_preference_feedback_component( + study_id: int, + storage: BaseStorage, + component_type: OUTPUT_COMPONENT_TYPE, + artifact_key: str | None = None, +) -> None: + value: dict[str, Any] = {"type": component_type} + if artifact_key is not None: + value["artifact_key"] = artifact_key + storage.set_study_system_attr( + study_id=study_id, + key=_SYSTEM_ATTR_FEEDBACK_COMPONENT, + value=value, + ) + + +def register_preference_feedback_component( + study: PreferentialStudy, + component_type: OUTPUT_COMPONENT_TYPE, + artifact_key: str | None = None, +) -> None: + """Register a preference feedback component to the study. + + With this feature, you can change the component, displayed on the + human feedback pages. By default, the Markdown note (``component_type="note"``) + is displayed. If you specify ``component_type="artifact"``, the viewer for the + specified artifact file will be displayed. + Args: + study: + The study to register the preference feedback component. + component_type: + The component type, displayed on the human feedback pages + (default: ``"note"``). + user_attr_artifact_key: + This option is required when the ``component_type`` is ``"artifact"``. + The user attribute, which is specified this field, must contain the + ``artifact``id you want to display on the human feedback page. + """ + if component_type == "artifact": + assert ( + artifact_key is not None + ), "artifact_key must be specified when component_type is Artifact" + + _register_preference_feedback_component( + study_id=study._study._study_id, + storage=study._study._storage, + component_type=component_type, + artifact_key=artifact_key, + ) diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index d3cd786f..5829f35d 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -15,6 +15,7 @@ from optuna.trial import FrozenTrial from . import _note as note from ._form_widget import get_form_widgets_json from ._named_objectives import get_objective_names +from ._preference_setting import _SYSTEM_ATTR_FEEDBACK_COMPONENT from ._preferential_history import _SYSTEM_ATTR_PREFIX_HISTORY from .artifact._backend import list_trial_artifacts from .preferential._study import _SYSTEM_ATTR_PREFERENTIAL_STUDY @@ -165,6 +166,9 @@ def serialize_study_detail( if form_widgets: serialized["form_widgets"] = form_widgets if serialized["is_preferential"]: + serialized["feedback_component_type"] = system_attrs.get( + _SYSTEM_ATTR_FEEDBACK_COMPONENT, {} + ) serialized["preference_history"] = serialize_preference_history(system_attrs) serialized["preferences"] = get_preferences(system_attrs) serialized["skipped_trials"] = skipped_trials diff --git a/python_tests/test_api.py b/python_tests/test_api.py index afe5a31c..db79c0d8 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -8,6 +8,7 @@ from optuna import get_all_study_summaries from optuna.study import StudyDirection from optuna_dashboard._app import create_app from optuna_dashboard._app import create_new_study +from optuna_dashboard._preference_setting import register_preference_feedback_component from optuna_dashboard._preferential_history import NewHistory from optuna_dashboard._preferential_history import remove_history from optuna_dashboard._preferential_history import report_history @@ -183,6 +184,36 @@ class APITestCase(TestCase): ) self.assertEqual(status, 400) + def test_change_component(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = create_study(storage=storage, n_generate=3) + register_preference_feedback_component(study, "note") + for _ in range(3): + study.ask() + + app = create_app(storage) + study_id = study._study._study_id + status, _, _ = send_request( + app, + f"/api/studies/{study_id}/preference_feedback_component_type", + "PUT", + body=json.dumps({"type": "artifact", "artifact_key": "image"}), + content_type="application/json", + ) + self.assertEqual(status, 204) + + status, _, body = send_request( + app, + f"/api/studies/{study_id}", + "GET", + content_type="application/json", + ) + self.assertEqual(status, 200) + + study_detail = json.loads(body) + assert study_detail["feedback_component_type"]["type"] == "artifact" + assert study_detail["feedback_component_type"]["artifact_key"] == "image" + def test_skip_trial(self) -> None: storage = optuna.storages.InMemoryStorage() study = create_study(n_generate=4, storage=storage) diff --git a/python_tests/test_preference_setting.py b/python_tests/test_preference_setting.py new file mode 100644 index 00000000..a0adcb39 --- /dev/null +++ b/python_tests/test_preference_setting.py @@ -0,0 +1,20 @@ +from __future__ import annotations + +from unittest import TestCase + +import optuna +from optuna_dashboard._preference_setting import _SYSTEM_ATTR_FEEDBACK_COMPONENT +from optuna_dashboard._preference_setting import register_preference_feedback_component +from optuna_dashboard.preferential._study import PreferentialStudy + + +class FeedbackSettingTestCase(TestCase): + def test_widget_to_dict_from_dict(self) -> None: + study = PreferentialStudy(optuna.create_study()) + register_preference_feedback_component(study, "artifact", "image_key") + system_attrs = study._study.system_attrs + feedback_type = system_attrs.get(_SYSTEM_ATTR_FEEDBACK_COMPONENT, {}) + assert "type" in feedback_type + assert feedback_type["type"] == "artifact" + assert "artifact_key" in feedback_type + assert feedback_type["artifact_key"] == "image_key"