From 5fcc64cff3e44bb34fb7ab9f2296fd7a8b4ff4e4 Mon Sep 17 00:00:00 2001 From: moririn2528 Date: Tue, 12 Sep 2023 16:00:46 +0900 Subject: [PATCH] fix by review --- optuna_dashboard/_app.py | 10 +++--- optuna_dashboard/_preference_setting.py | 24 ++++++-------- optuna_dashboard/_serializer.py | 10 +++--- python_tests/test_api.py | 44 ++++--------------------- python_tests/test_preference_setting.py | 14 ++++---- 5 files changed, 34 insertions(+), 68 deletions(-) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index fad4b74c..2ad7874e 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -28,7 +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_output_component +from ._preference_setting import _register_preference_feedback_component_type from ._preferential_history import NewHistory from ._preferential_history import report_history from ._rdb_migration import register_rdb_migration_route @@ -307,11 +307,11 @@ def create_app( response.status = 204 return {} - @app.post("/api/studies//component") + @app.put("/api/studies//preference_feedback_component_type") @json_api_view - def post_component(study_id: int) -> dict[str, Any]: + def put_component(study_id: int) -> dict[str, Any]: try: - component_type = request.json.get("component_type", "") + component_type = request.json.get("type", "") artifact_key = request.json.get("artifact_key", None) except ValueError: response.status = 400 @@ -320,7 +320,7 @@ def create_app( response.status = 400 return {"reason": "component_type must be either 'Note' or 'Artifact'."} - _register_output_component( + _register_preference_feedback_component_type( study_id=study_id, storage=storage, component_type=component_type, diff --git a/optuna_dashboard/_preference_setting.py b/optuna_dashboard/_preference_setting.py index 93b1a34c..73e9b489 100644 --- a/optuna_dashboard/_preference_setting.py +++ b/optuna_dashboard/_preference_setting.py @@ -12,30 +12,26 @@ if TYPE_CHECKING: OUTPUT_COMPONENT_TYPE = Literal["Note", "Artifact"] -_SYSTEM_ATTR_FEEDBACK_COMPONENT_TYPE = "preference:component_type" -_SYSTEM_ATTR_FEEDBACK_ARTIFACT_KEY = "preference:component_artifact_key" +_SYSTEM_ATTR_FEEDBACK_COMPONENT = "preference:component" -def _register_output_component( +def _register_preference_feedback_component_type( study_id: int, storage: BaseStorage, component_type: OUTPUT_COMPONENT_TYPE, - artifact_key: str | None = None, + artifact_key: str = "", ) -> None: storage.set_study_system_attr( study_id=study_id, - key=_SYSTEM_ATTR_FEEDBACK_COMPONENT_TYPE, - value=component_type, + key=_SYSTEM_ATTR_FEEDBACK_COMPONENT, + value={ + "type": component_type, + "artifact_key": artifact_key, + } ) - if artifact_key is not None: - storage.set_study_system_attr( - study_id=study_id, - key=_SYSTEM_ATTR_FEEDBACK_ARTIFACT_KEY, - value=artifact_key, - ) -def register_output_component( +def register_preference_feedback_component_type( study: PreferentialStudy, component_type: OUTPUT_COMPONENT_TYPE, artifact_key: str = "", @@ -52,7 +48,7 @@ def register_output_component( this argument is used as the attribute key of the artifact. Each trial displays the artifact whose id is the value of the attribute. """ - _register_output_component( + _register_preference_feedback_component_type( study_id=study._study._study_id, storage=study._study._storage, component_type=component_type, diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 511e7365..98db5ffb 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -15,8 +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_ARTIFACT_KEY -from ._preference_setting import _SYSTEM_ATTR_FEEDBACK_COMPONENT_TYPE +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 @@ -164,10 +163,9 @@ def serialize_study_detail( form_widgets = get_form_widgets_json(system_attrs) if form_widgets: serialized["form_widgets"] = form_widgets - if _SYSTEM_ATTR_FEEDBACK_COMPONENT_TYPE in system_attrs: - serialized["feedback_component_type"] = system_attrs[_SYSTEM_ATTR_FEEDBACK_COMPONENT_TYPE] - if _SYSTEM_ATTR_FEEDBACK_ARTIFACT_KEY in system_attrs: - serialized["feedback_artifact_key"] = system_attrs[_SYSTEM_ATTR_FEEDBACK_ARTIFACT_KEY] + if serialized["is_preferential"]: + serialized["feedback_component_type"] = system_attrs.get( + _SYSTEM_ATTR_FEEDBACK_COMPONENT, {}) if serialized["is_preferential"]: serialized["preference_history"] = serialize_preference_history(system_attrs) serialized["preferences"] = get_preferences(system_attrs) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index e22d104c..1908c188 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -8,7 +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_output_component +from optuna_dashboard._preference_setting import register_preference_feedback_component_type from optuna_dashboard.preferential import create_study from .wsgi_client import send_request @@ -183,7 +183,7 @@ class APITestCase(TestCase): def test_change_component(self) -> None: storage = optuna.storages.InMemoryStorage() study = create_study(storage=storage, n_generate=3) - register_output_component(study, "Note") + register_preference_feedback_component_type(study, "Note") for _ in range(3): study.ask() @@ -191,9 +191,9 @@ class APITestCase(TestCase): study_id = study._study._study_id status, _, _ = send_request( app, - f"/api/studies/{study_id}/component", - "POST", - body=json.dumps({"component_type": "Artifact", "artifact_key": "image"}), + 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) @@ -207,38 +207,8 @@ class APITestCase(TestCase): self.assertEqual(status, 200) study_detail = json.loads(body) - assert study_detail["feedback_component_type"] == "Artifact" - assert study_detail["feedback_artifact_key"] == "image" - - def test_change_component_type_only(self) -> None: - storage = optuna.storages.InMemoryStorage() - study = create_study(storage=storage, n_generate=3) - register_output_component(study, "Artifact", "audio") - 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}/component", - "POST", - body=json.dumps({"component_type": "Note"}), - 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"] == "Note" - assert study_detail["feedback_artifact_key"] == "audio" + 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() diff --git a/python_tests/test_preference_setting.py b/python_tests/test_preference_setting.py index 033a9092..8a601af7 100644 --- a/python_tests/test_preference_setting.py +++ b/python_tests/test_preference_setting.py @@ -3,16 +3,18 @@ from __future__ import annotations from unittest import TestCase import optuna -from optuna_dashboard._preference_setting import _SYSTEM_ATTR_FEEDBACK_ARTIFACT_KEY -from optuna_dashboard._preference_setting import _SYSTEM_ATTR_FEEDBACK_COMPONENT_TYPE -from optuna_dashboard._preference_setting import register_output_component +from optuna_dashboard._preference_setting import _SYSTEM_ATTR_FEEDBACK_COMPONENT +from optuna_dashboard._preference_setting import register_preference_feedback_component_type 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_output_component(study, "Artifact", "image_key") + register_preference_feedback_component_type(study, "Artifact", "image_key") system_attrs = study._study.system_attrs - assert system_attrs.get(_SYSTEM_ATTR_FEEDBACK_COMPONENT_TYPE, "") == "Artifact" - assert system_attrs.get(_SYSTEM_ATTR_FEEDBACK_ARTIFACT_KEY, "") == "image_key" + 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"