From f8390ee2dd279517106d8ceebc559e36ae7b7b2f Mon Sep 17 00:00:00 2001 From: porink0424 Date: Fri, 9 Aug 2024 11:35:06 +0900 Subject: [PATCH 1/2] Add set_user_attr in rename_study api --- optuna_dashboard/_app.py | 2 ++ 1 file changed, 2 insertions(+) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index bdbfb168..a983d320 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -159,6 +159,8 @@ def create_app( storage=storage, study_name=dst_study_name, directions=src_study.directions ) dst_study.add_trials(src_study.get_trials(deepcopy=False)) + for key, value in src_study.user_attrs.items(): + dst_study.set_user_attr(key, value) note.copy_notes(storage, src_study, dst_study) except DuplicatedStudyError: response.status = 400 # Bad request From 46d664ed7272f95049c933524e9ae26695b880e3 Mon Sep 17 00:00:00 2001 From: porink0424 Date: Fri, 9 Aug 2024 14:09:48 +0900 Subject: [PATCH 2/2] Add a test case --- python_tests/test_api.py | 22 ++++++++++++++++++++++ 1 file changed, 22 insertions(+) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 818e4465..d6732903 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -726,6 +726,28 @@ class APITestCase(TestCase): ) self.assertEqual(status, 400) + def test_rename_study(self) -> None: + storage = optuna.storages.InMemoryStorage() + study = optuna.create_study(study_name="foo", storage=storage) + study.set_user_attr("key1", "value1") + study.optimize(objective, n_trials=2) + + app = create_app(storage) + status, _, _ = send_request( + app, + f"/api/studies/{study._study_id}/rename", + "POST", + body=json.dumps({"study_name": "bar"}), + content_type="application/json", + ) + self.assertEqual(status, 201) + + renamed_study = optuna.load_study(study_name="bar", storage=storage) + self.assertNotEqual(study._study_id, renamed_study._study_id) + self.assertEqual(len(renamed_study.trials), 2) + self.assertEqual(renamed_study.user_attrs, {"key1": "value1"}) + self.assertEqual(len(get_all_study_summaries(storage)), 1) + class BottleRequestHookTestCase(TestCase): def test_ignore_trailing_slashes(self) -> None: