diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index d7ddf72b..1baf76fe 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -30,7 +30,10 @@ 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_type from ._preferential_history import NewHistory +from ._preferential_history import PreferenceHistoryNotFound +from ._preferential_history import remove_history from ._preferential_history import report_history +from ._preferential_history import restore_history from ._rdb_migration import register_rdb_migration_route from ._serializer import serialize_study_detail from ._serializer import serialize_study_summary @@ -326,6 +329,28 @@ def create_app( component_type=component_type, artifact_key=artifact_key, ) + + @app.delete("/api/studies//preference/") + @json_api_view + def remove_preference(study_id: int, history_id: str) -> dict[str, Any]: + try: + remove_history(study_id, storage, history_id) + except PreferenceHistoryNotFound: + response.status = 404 + return {"reason": f"history_id={history_id} is not found"} + + response.status = 204 + return {} + + @app.post("/api/studies//preference/") + @json_api_view + def restore_preference(study_id: int, history_id: str) -> dict[str, Any]: + try: + restore_history(study_id, storage, history_id) + except PreferenceHistoryNotFound: + response.status = 404 + return {"reason": f"history_id={history_id} is not found"} + response.status = 204 return {} diff --git a/optuna_dashboard/_preferential_history.py b/optuna_dashboard/_preferential_history.py index 6b81b9bb..ef3f87f4 100644 --- a/optuna_dashboard/_preferential_history.py +++ b/optuna_dashboard/_preferential_history.py @@ -4,10 +4,10 @@ from dataclasses import dataclass from datetime import datetime import json from typing import TYPE_CHECKING -import uuid from optuna.storages import BaseStorage +from .preferential._system_attrs import _SYSTEM_ATTR_PREFIX_PREFERENCE from .preferential._system_attrs import report_preferences @@ -23,13 +23,24 @@ if TYPE_CHECKING: { "mode": FeedbackMode, "id": str, - "preference_id": str, "timestamp": str, "candidates": list[int], "clicked": int, + "preferences": list[tuple[int, int]], }, ) History = ChooseWorstHistory + SerializedHistory = TypedDict( + "SerializedHistory", + { + "history": History, + "is_removed": bool, + }, + ) + + +class PreferenceHistoryNotFound(Exception): + pass @dataclass @@ -43,14 +54,14 @@ def report_history( study_id: int, storage: BaseStorage, input_data: NewHistory, -) -> None: +) -> str: preferences = [] # TODO(moririn): Use TypeGuard after adding other history types. if input_data.mode == "ChooseWorst": preferences = [ - (best, input_data.clicked) - for best in input_data.candidates - if best != input_data.clicked + (better, input_data.clicked) + for better in input_data.candidates + if better != input_data.clicked ] else: assert False, f"Unknown data: {input_data}" @@ -60,21 +71,40 @@ def report_history( storage=storage, preferences=preferences, ) - history_id = str(uuid.uuid4()) if input_data.mode == "ChooseWorst": history: ChooseWorstHistory = { "mode": "ChooseWorst", - "id": history_id, - "preference_id": preference_id, + "id": preference_id, "timestamp": datetime.now().isoformat(), "candidates": input_data.candidates, "clicked": input_data.clicked, + "preferences": preferences, } - key = _SYSTEM_ATTR_PREFIX_HISTORY + history_id + key = _SYSTEM_ATTR_PREFIX_HISTORY + preference_id storage.set_study_system_attr( study_id=study_id, key=key, value=json.dumps(history), ) + return preference_id + + +def remove_history(study_id: int, storage: BaseStorage, history_id: str) -> None: + system_attrs = storage.get_study_system_attrs(study_id) + history_key = _SYSTEM_ATTR_PREFIX_HISTORY + history_id + if history_key not in system_attrs: + raise PreferenceHistoryNotFound + storage.set_study_system_attr(study_id, _SYSTEM_ATTR_PREFIX_PREFERENCE + history_id, []) + + +def restore_history(study_id: int, storage: BaseStorage, history_id: str) -> None: + system_attrs = storage.get_study_system_attrs(study_id) + history_key = _SYSTEM_ATTR_PREFIX_HISTORY + history_id + if history_key not in system_attrs: + raise PreferenceHistoryNotFound + history: History = json.loads(system_attrs.get(history_key, "")) + storage.set_study_system_attr( + study_id, _SYSTEM_ATTR_PREFIX_PREFERENCE + history_id, history["preferences"] + ) diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index f9aeabec..ea22644f 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -20,14 +20,15 @@ from ._preferential_history import _SYSTEM_ATTR_PREFIX_HISTORY from .artifact._backend import list_trial_artifacts from .preferential._study import _SYSTEM_ATTR_PREFERENTIAL_STUDY from .preferential._system_attrs import get_preferences +from .preferential._system_attrs import is_preference_removed if TYPE_CHECKING: from typing import Literal from typing import TypedDict - from ._preferential_history import ChooseWorstHistory from ._preferential_history import History + from ._preferential_history import SerializedHistory Attribute = TypedDict( "Attribute", @@ -179,24 +180,29 @@ def serialize_study_detail( def serialize_preference_history( system_attrs: dict[str, Any], -) -> list[History]: - histories: list[History] = [] +) -> list[SerializedHistory]: + histories: list[SerializedHistory] = [] for k, v in system_attrs.items(): if not k.startswith(_SYSTEM_ATTR_PREFIX_HISTORY): continue choice: dict[str, Any] = json.loads(v) if choice["mode"] == "ChooseWorst": - history: ChooseWorstHistory = { + history: History = { "mode": "ChooseWorst", "id": choice["id"], - "preference_id": choice["preference_id"], "timestamp": choice["timestamp"], "candidates": choice["candidates"], "clicked": choice["clicked"], + "preferences": choice["preferences"], } - histories.append(history) + histories.append( + { + "history": history, + "is_removed": is_preference_removed(system_attrs, choice["id"]), + } + ) - histories.sort(key=lambda c: datetime.fromisoformat(c["timestamp"])) + histories.sort(key=lambda c: datetime.fromisoformat(c["history"]["timestamp"])) return histories diff --git a/optuna_dashboard/preferential/_system_attrs.py b/optuna_dashboard/preferential/_system_attrs.py index 47c2a486..438b411c 100644 --- a/optuna_dashboard/preferential/_system_attrs.py +++ b/optuna_dashboard/preferential/_system_attrs.py @@ -44,6 +44,12 @@ def get_preferences(study_system_attrs: dict[str, Any]) -> list[tuple[int, int]] return preferences +def is_preference_removed(study_system_attrs: dict[str, Any], preference_id: str) -> bool: + key = _SYSTEM_ATTR_PREFIX_PREFERENCE + preference_id + preference = study_system_attrs.get(key, []) + return len(preference) == 0 + + def report_skip( study_id: int, trial_id: int, diff --git a/python_tests/test_api.py b/python_tests/test_api.py index c478dc26..2e46be43 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -8,7 +8,11 @@ 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._preferential_history import NewHistory from optuna_dashboard._preference_setting import register_preference_feedback_component_type +from optuna_dashboard._preferential_history import remove_history +from optuna_dashboard._preferential_history import report_history +from optuna_dashboard._serializer import serialize_preference_history from optuna_dashboard.preferential import create_study from .wsgi_client import send_request @@ -234,6 +238,82 @@ class APITestCase(TestCase): assert len(best_trials) == 1 assert best_trials[0].number == 2 + def test_remove_history(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = create_study(storage=storage, n_generate=3) + for _ in range(3): + study.ask() + + app = create_app(storage) + study_id = study._study._study_id + history_id = report_history( + study_id, + storage, + NewHistory( + mode="ChooseWorst", + candidates=[0, 1, 2], + clicked=2, + ), + ) + histories = serialize_preference_history(storage.get_study_system_attrs(study_id)) + assert len(histories) == 1 + assert not histories[0]["is_removed"] + + status, _, _ = send_request( + app, + f"/api/studies/{study_id}/preference/{history_id}", + "DELETE", + content_type="application/json", + ) + self.assertEqual(status, 204) + histories = serialize_preference_history(storage.get_study_system_attrs(study_id)) + assert len(histories) == 1 + assert histories[0]["is_removed"] + assert len(study.get_preferences()) == 0 + + def test_restore_history(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = create_study(storage=storage, n_generate=3) + for _ in range(3): + study.ask() + + app = create_app(storage) + study_id = study._study._study_id + history_id = report_history( + study_id, + storage, + NewHistory( + mode="ChooseWorst", + candidates=[0, 1, 2], + clicked=2, + ), + ) + remove_history(study_id, storage, history_id) + histories = serialize_preference_history(storage.get_study_system_attrs(study_id)) + assert len(histories) == 1 + assert histories[0]["is_removed"] + assert len(study.get_preferences()) == 0 + + status, _, _ = send_request( + app, + f"/api/studies/{study_id}/preference/{history_id}", + "POST", + content_type="application/json", + ) + self.assertEqual(status, 204) + histories = serialize_preference_history(storage.get_study_system_attrs(study_id)) + assert len(histories) == 1 + assert not histories[0]["is_removed"] + preferences = study.get_preferences() + preferences.sort(key=lambda x: (x[0].number, x[1].number)) + assert len(preferences) == 2 + better, worse = preferences[0] + assert better.number == 0 + assert worse.number == 2 + better, worse = preferences[1] + assert better.number == 1 + assert worse.number == 2 + def test_create_study(self) -> None: for name, directions, expected_status in [ ("single-objective success", ["minimize"], 201), diff --git a/python_tests/test_preferential_history.py b/python_tests/test_preferential_history.py index 51c9b0f8..3d1a7d70 100644 --- a/python_tests/test_preferential_history.py +++ b/python_tests/test_preferential_history.py @@ -1,9 +1,15 @@ from __future__ import annotations +import json from typing import Callable +from typing import TYPE_CHECKING +from optuna.storages import BaseStorage +from optuna_dashboard._preferential_history import _SYSTEM_ATTR_PREFIX_HISTORY from optuna_dashboard._preferential_history import NewHistory +from optuna_dashboard._preferential_history import remove_history from optuna_dashboard._preferential_history import report_history +from optuna_dashboard._preferential_history import restore_history from optuna_dashboard._serializer import serialize_preference_history from optuna_dashboard.preferential import create_study from optuna_dashboard.preferential._system_attrs import _SYSTEM_ATTR_PREFIX_PREFERENCE @@ -12,6 +18,10 @@ from .storage_supplier import parametrize_storages from .storage_supplier import StorageSupplier +if TYPE_CHECKING: + from optuna_dashboard._preferential_history import History + + @parametrize_storages def test_report_and_get_choices(storage_supplier: Callable[[], StorageSupplier]) -> None: with storage_supplier() as storage: @@ -25,37 +35,102 @@ def test_report_and_get_choices(storage_supplier: Callable[[], StorageSupplier]) report_history( study_id=study_id, storage=storage, - input_data=NewHistory( - mode="ChooseWorst", - candidates=[0, 1, 2], - clicked=1, - ), + input_data=NewHistory(mode="ChooseWorst", candidates=[0, 1, 2], clicked=1), ) report_history( study_id=study_id, storage=storage, - input_data=NewHistory( - mode="ChooseWorst", - candidates=[0, 2, 3, 4], - clicked=0, - ), + input_data=NewHistory(mode="ChooseWorst", candidates=[0, 2, 3, 4], clicked=0), ) history = serialize_preference_history(storage.get_study_system_attrs(study_id)) sys_attrs = storage.get_study_system_attrs(study_id) assert len(history) == 2 - assert history[0]["candidates"] == [0, 1, 2] - assert history[0]["clicked"] == 1 - preferences = sys_attrs[_SYSTEM_ATTR_PREFIX_PREFERENCE + history[0]["preference_id"]] + assert history[0]["history"]["candidates"] == [0, 1, 2] + assert history[0]["history"]["clicked"] == 1 + preferences = sys_attrs[_SYSTEM_ATTR_PREFIX_PREFERENCE + history[0]["history"]["id"]] assert len(preferences) == 2 for i, (best, worst) in enumerate([(0, 1), (2, 1)]): assert len(preferences[i]) == 2 assert preferences[i][0] == best assert preferences[i][1] == worst - assert history[1]["candidates"] == [0, 2, 3, 4] - assert history[1]["clicked"] == 0 - preferences = sys_attrs[_SYSTEM_ATTR_PREFIX_PREFERENCE + history[1]["preference_id"]] + assert history[1]["history"]["candidates"] == [0, 2, 3, 4] + assert history[1]["history"]["clicked"] == 0 + preferences = sys_attrs[_SYSTEM_ATTR_PREFIX_PREFERENCE + history[1]["history"]["id"]] assert len(preferences) == 3 for i, (best, worst) in enumerate([(2, 0), (3, 0), (4, 0)]): assert len(preferences[i]) == 2 assert preferences[i][0] == best assert preferences[i][1] == worst + + +def get_preferences_history( + study_id: int, + storage: BaseStorage, + history_id: str, +) -> tuple[list[tuple[int, int]], History]: + system_attrs = storage.get_study_system_attrs(study_id) + history: History = json.loads(system_attrs.get(_SYSTEM_ATTR_PREFIX_HISTORY + history_id, "")) + preference: list[tuple[int, int]] = system_attrs.get( + _SYSTEM_ATTR_PREFIX_PREFERENCE + history_id, [] + ) + return preference, history + + +@parametrize_storages +def test_remove_history(storage_supplier: Callable[[], StorageSupplier]) -> None: + with storage_supplier() as storage: + study = create_study(storage=storage, n_generate=5) + for _ in range(5): + trial = study.ask() + trial.suggest_float("x", 0, 1) + study_id = study._study._study_id + + history_id = report_history( + study_id=study_id, + storage=storage, + input_data=NewHistory(mode="ChooseWorst", candidates=[0, 1, 2], clicked=1), + ) + remove_history(study_id, storage, history_id) + preference, history = get_preferences_history(study_id, storage, history_id) + assert history["mode"] == "ChooseWorst" + assert history["candidates"] == [0, 1, 2] + assert history["clicked"] == 1 + assert len(preference) == 0 + + remove_history(study_id, storage, history_id) + preference, history = get_preferences_history(study_id, storage, history_id) + assert len(preference) == 0 + + +@parametrize_storages +def test_restore_history(storage_supplier: Callable[[], StorageSupplier]) -> None: + with storage_supplier() as storage: + study = create_study(storage=storage, n_generate=5) + for _ in range(5): + trial = study.ask() + trial.suggest_float("x", 0, 1) + study_id = study._study._study_id + + history_id = report_history( + study_id=study_id, + storage=storage, + input_data=NewHistory(mode="ChooseWorst", candidates=[0, 1, 2], clicked=1), + ) + remove_history(study_id, storage, history_id) + preference, history = get_preferences_history(study_id, storage, history_id) + assert len(preference) == 0 + + restore_history(study_id, storage, history_id) + preference, history = get_preferences_history(study_id, storage, history_id) + assert history["mode"] == "ChooseWorst" + assert history["candidates"] == [0, 1, 2] + assert history["clicked"] == 1 + assert len(preference) == 2 + for i, (best, worst) in enumerate([(0, 1), (2, 1)]): + assert len(preference[i]) == 2 + assert preference[i][0] == best + assert preference[i][1] == worst + + restore_history(study_id, storage, history_id) + preference, history = get_preferences_history(study_id, storage, history_id) + assert len(preference) == 2