From 2965e69508d4747d0b5bf9a43575fd498acbab1d Mon Sep 17 00:00:00 2001 From: Victoria A <52001888+adjeiv@users.noreply.github.com> Date: Sun, 12 Nov 2023 10:49:51 +0000 Subject: [PATCH] Retain trial notes --- optuna_dashboard/_app.py | 4 +++- optuna_dashboard/_note.py | 40 +++++++++++++++++++++++++++++++++----- python_tests/test_note.py | 41 +++++++++++++++++++++++++++++++++++++++ 3 files changed, 79 insertions(+), 6 deletions(-) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 812201ad..9ce820b7 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -164,7 +164,8 @@ def create_app( if new_study_summary is None: response.status = 500 return {"reason": "Failed to load the new study"} - + + note.transfer_notes(storage, src_study, dst_study) storage.delete_study(src_study._study_id) response.status = 201 return serialize_study_summary(new_study_summary) @@ -176,6 +177,7 @@ def create_app( delete_all_artifacts(artifact_store, storage, study_id) try: + note.delete_study_notes(storage, study_id) storage.delete_study(study_id) except KeyError: response.status = 404 # Not found diff --git a/optuna_dashboard/_note.py b/optuna_dashboard/_note.py index 95853a3d..e510ddfd 100644 --- a/optuna_dashboard/_note.py +++ b/optuna_dashboard/_note.py @@ -109,6 +109,21 @@ def note_str_key_prefix(trial_id: Optional[int]) -> str: return prefix return f"dashboard:{trial_id}:note_str:" +def transfer_notes(storage: BaseStorage, src_study: optuna.Study, dst_study: optuna.Study) -> None: + system_attrs = storage.get_study_system_attrs(study_id=src_study._study_id) + + def transfer(src_trial_id: Optional[int], dst_trial_id: Optional[int]) -> None: + note = get_note_from_system_attrs(system_attrs, src_trial_id)["body"] + save_note_with_version(storage, dst_study._study_id, dst_trial_id, 0, note) + delete_notes(storage, src_study._study_id, src_trial_id) + + # Transfer individual trial notes + for src_trial, dst_trial in zip(src_study.get_trials(), dst_study.get_trials()): + transfer(src_trial._trial_id, dst_trial._trial_id) + + # Transfer study note + NO_SRC_TRIAL, NO_DST_TRIAL = None, None + transfer(NO_SRC_TRIAL, NO_DST_TRIAL) def get_note_from_system_attrs(system_attrs: dict[str, Any], trial_id: Optional[int]) -> NoteType: if note_ver_key(trial_id) not in system_attrs: @@ -131,6 +146,13 @@ def version_is_incremented( db_note_ver = system_attrs.get(note_ver_key(trial_id), 0) return req_note_ver == db_note_ver + 1 +def all_trial_notes(storage: BaseStorage, study_id: int, trial_id: Optional[int]) -> dict[str, str]: + all_note_attrs: dict[str, str] = { + key: value + for key, value in storage.get_study_system_attrs(study_id).items() + if key.startswith(note_str_key_prefix(trial_id)) + } + return all_note_attrs def save_note_with_version( storage: BaseStorage, study_id: int, trial_id: Optional[int], ver: int, body: str @@ -142,15 +164,23 @@ def save_note_with_version( storage.set_study_system_attr(study_id, k, v) # Clear previous messages - all_note_attrs: dict[str, str] = { - key: value - for key, value in storage.get_study_system_attrs(study_id).items() - if key.startswith(note_str_key_prefix(trial_id)) - } + all_note_attrs = all_trial_notes(storage, study_id, trial_id) if len(all_note_attrs) > len(attrs): for i in range(len(attrs), len(all_note_attrs)): storage.set_study_system_attr(study_id, f"{note_str_key_prefix(trial_id)}{i}", "") +def delete_study_notes(storage: BaseStorage, study_id): + study = storage.get_study_name_from_id(study_id) + for trial in storage.get_all_trials(study_id): + delete_notes(storage, study_id, trial._trial_id) + + delete_notes(storage, study_id, None) + +def delete_notes(storage: BaseStorage, study_id: int, trial_id: Optional[int]) -> None: + all_note_attrs = all_trial_notes(storage, study_id, trial_id) + + for i in range(len(all_note_attrs)): + storage.set_study_system_attr(study_id, f"{note_str_key_prefix(trial_id)}{i}", "") def split_body(note_str: str, trial_id: Optional[int]) -> dict[str, str]: note_len = len(note_str) diff --git a/python_tests/test_note.py b/python_tests/test_note.py index 7cc02cd4..4997d16a 100644 --- a/python_tests/test_note.py +++ b/python_tests/test_note.py @@ -53,3 +53,44 @@ class NoteTestCase(TestCase): note_dict = note.get_note_from_system_attrs(system_attrs, trial._trial_id) self.assertEqual(note_dict["body"], body) self.assertEqual(note_dict["version"], expected_ver) + + def test_delete_notes_trial(self) -> None: + study = optuna.create_study() + trial_1 = study.ask({"x1": optuna.distributions.FloatDistribution(0, 10)}) + trial_2 = study.ask({"x1": optuna.distributions.FloatDistribution(0, 10)}) + storage = study._storage + + for trial, body in [(trial_1, "version 1"), (trial_2, "version 2")]: + save_note(trial, body) + + # first assert existence + actual = get_note(trial) + self.assertEqual(actual, body) + + # delete + note.delete_notes(storage, study._study_id, trial._trial_id) + + # assert deletion + actual = get_note(trial) + self.assertEqual(actual, "") + + + note.delete_study_notes(storage, study._study_id) + + def test_delete_notes_study(self) -> None: + pass + + def test_transfer_notes(self) -> None: + study = optuna.create_study() + trial_1 = study.ask({"x1": optuna.distributions.FloatDistribution(0, 10)}) + trial_2 = study.ask({"x2": optuna.distributions.FloatDistribution(0, 10)}) + storage = study._storage + + save_note(trial_1, "trial 1") + save_note(trial_2, "trial 2") + + new_study = optuna.create_study( + storage=storage, directions=study.directions + ) + note.transfer_notes(storage, study, new_study) +