From ad8c1531b9da4d820aea06e4b8d643b42e21748e Mon Sep 17 00:00:00 2001 From: c-bata Date: Mon, 14 Aug 2023 19:58:54 +0900 Subject: [PATCH 01/31] Write docstring for preferential optimization --- docs/api.rst | 14 ++ optuna_dashboard/preferential/_study.py | 230 ++++++++++++++++++++++++ 2 files changed, 244 insertions(+) diff --git a/docs/api.rst b/docs/api.rst index f1a0b3c6..09a9c14b 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -18,6 +18,9 @@ General APIs Human-in-the-loop ----------------- +Form Widgets +~~~~~~~~~~~~ + .. autosummary:: :toctree: _generated/ :nosignatures: @@ -30,6 +33,17 @@ Human-in-the-loop optuna_dashboard.TextInputWidget optuna_dashboard.ObjectiveUserAttrRef +Preferential Optimization +~~~~~~~~~~~~~~~~~~~~~~~~~ + +.. autosummary:: + :toctree: _generated/ + :nosignatures: + + optuna_dashboard.preferential.create_study + optuna_dashboard.preferential.load_study + optuna_dashboard.preferential.PreferentialStudy + Streamlit ----------------- diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index 9bbd6d2f..029a04e0 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -22,15 +22,49 @@ _SYSTEM_ATTR_COMPARISON_READY = "preference:comparison_ready" class PreferentialStudy: + """A Study-like class for preferential optimization. + + This object provides interfaces to create a new `Trial`_, set/get results + of pairwise comparison called preferences. + + .. _Trial: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.trial.Trial.html#optuna.trial.Trial + + Note that the direct use of this constructor is not recommended. + To create and load a study, please refer to the documentation of + :func:`~optuna_dashboard.preferential.create_study` and + :func:`~optuna_dashboard.preferential.load_study` respectively. + """ def __init__(self, study: optuna.Study) -> None: self._study = study @property def trials(self) -> list[FrozenTrial]: + """Return the all trials. + + .. seealso:: + + See `Study.trials`_ for details. + + .. _Study.trials: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.trials + + Returns: + A list of FrozenTrial object + """ return self._study.trials @property def best_trials(self) -> list[FrozenTrial]: + """Return the trials that is not dominated by other trials. + + .. seealso:: + + See `Study.best_trials`_ for details. + + .. _Study.best_trials: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.best_trials + + Returns: + A list of FrozenTrial object + """ ready_trials = [ t for t in self._study.get_trials( @@ -44,14 +78,35 @@ class PreferentialStudy: @property def study_name(self) -> str: + """Return the name of the study. + + Returns: + A string object + """ return self._study.study_name @property def user_attrs(self) -> dict[str, Any]: + """Return user attributes of the study. + + .. seealso:: + + See `Study.user_attrs`_ for details. + + .. _Study.user_attrs: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.user_attrs + + Returns: + A dictionary containing all user attributes + """ return self._study.user_attrs @property def preferences(self) -> list[tuple[FrozenTrial, FrozenTrial]]: + """Return results of pairwise comparison. + + Returns: + A list of the pair of FrozenTrial objects. The left trial is better than the right one. + """ return self.get_preferences(deepcopy=True) def get_trials( @@ -59,15 +114,69 @@ class PreferentialStudy: deepcopy: bool = True, states: Container[optuna.trial.TrialState] | None = None, ) -> list[FrozenTrial]: + """Return the trials that is not dominated by other trials. + + .. seealso:: + + See `Study.get_trials`_ for details. + + .. _Study.get_trials: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.get_trials + + Args: + deepcopy: + Flag to control whether to apply ``copy.deepcopy()`` to the trials. + Note that if you set the flag to :obj:`False`, you shouldn't mutate + any fields of the returned trial. Otherwise the internal state of + the study may corrupt and unexpected behavior may happen. + states: + Trial states to filter on. If :obj:`None`, include all states. + + Returns: + A list of FrozenTrial object + """ return self._study.get_trials(deepcopy, states) def ask(self, fixed_distributions: dict[str, BaseDistribution] | None = None) -> optuna.Trial: + """Create a new trial from which hyperparameters can be suggested. + + .. seealso:: + + See `Study.ask`_ for details. + + .. _Study.ask: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.ask + + Args: + fixed_distributions: + A dictionary containing the parameter names and parameter's distributions. Each + parameter in this dictionary is automatically suggested for the returned trial, + even when the suggest method is not explicitly invoked by the user. If this + argument is set to :obj:`None`, no parameter is automatically suggested. + + Returns: + A Trial object. + """ return self._study.ask(fixed_distributions) def add_trial(self, trial: FrozenTrial) -> None: + """Add a trial to the study. + + .. seealso:: + + See `Study.add_trials()`_ for details. + + .. _Study.add_trials(): https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.add_trials + """ self._study.add_trial(trial) def add_trials(self, trials: Iterable[FrozenTrial]) -> None: + """Add trials to the study. + + .. seealso:: + + See `Study.add_trials()`_ for details. + + .. _Study.add_trials(): https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.add_trials + """ self._study.add_trials(trials) def report_preference( @@ -75,6 +184,14 @@ class PreferentialStudy: better_trials: FrozenTrial | list[FrozenTrial], worse_trials: FrozenTrial | list[FrozenTrial], ) -> None: + """Report results of pairwise comparison. + + Args: + better_trials: + Trials that are better than worse_trials. + worse_trials: + Trials that are worse than better_trials. + """ if not isinstance(better_trials, list): better_trials = [better_trials] if not isinstance(worse_trials, list): @@ -83,12 +200,41 @@ class PreferentialStudy: report_preferences(self._study, [(b, w) for b in better_trials for w in worse_trials]) def get_preferences(self, *, deepcopy: bool = True) -> list[tuple[FrozenTrial, FrozenTrial]]: + """Return results of pairwise comparison. + + Args: + deepcopy: + Flag to control whether to apply ``copy.deepcopy()`` to the trials. + Note that if you set the flag to :obj:`False`, you shouldn't mutate + any fields of the returned trial. Otherwise the internal state of + the study may corrupt and unexpected behavior may happen. + + Returns: + A list of the pair of FrozenTrial objects. The left trial is better than the right one. + """ return get_preferences(self._study, deepcopy=deepcopy) def set_user_attr(self, key: str, value: Any) -> None: + """Set a user attribute to the study. + + Args: + key: A key string of the attribute. + value: A value of the attribute. The value should be JSON serializable. + + .. seealso:: + + See the `tutorial for user attributes `_on Optuna's documentation. + """ self._study.set_user_attr(key, value) def mark_comparison_ready(self, trial_or_number: optuna.Trial | int) -> None: + """Mark trials ready to compare. + + Args: + trial_or_number: + A Trial object or trial_number. + """ storage = self._study._storage if isinstance(trial_or_number, optuna.Trial): trial_id = trial_or_number._trial_id @@ -108,6 +254,45 @@ def create_study( study_name: str | None = None, load_if_exists: bool = False, ) -> PreferentialStudy: + """Like ``optuna.create_study()``, but for preferential optimization. + + Example: + + .. testcode:: + + import optuna + from optuna_dashboard.preferential import create_study + + + study = create_study() + trial = study.ask() + + Args: + storage: + Database URL. If this argument is set to None, in-memory storage is used, and the + :class:`~optuna_dashboard.preferential.PreferentialStudy` will not be persistent. + + sampler: + A sampler object that implements background algorithm for value suggestion. + If :obj:`None` is specified, `RandomSampler`_ is used. Please note that + most Optuna samplers does not work efficiently for preferential optimization. + + .. _RandomSampler: https://optuna.readthedocs.io/en/stable/reference/samplers/generated/optuna.samplers.RandomSampler.html + + study_name: + Study's name. If this argument is set to None, a unique name is generated + automatically. + + load_if_exists: + Flag to control the behavior to handle a conflict of study names. + In the case where a study named ``study_name`` already exists in the ``storage``, + a :class:`~optuna.exceptions.DuplicatedStudyError` is raised if ``load_if_exists`` is + set to :obj:`False`. + Otherwise, the creation of the study is skipped, and the existing one is returned. + + Returns: + A :class:`~optuna_dashboard.preferential.PreferentialStudy` object. + """ try: study = optuna.create_study( storage=storage, @@ -143,6 +328,51 @@ def load_study( storage: str | optuna.storages.BaseStorage, sampler: BaseSampler | None = None, ) -> PreferentialStudy: + """Like ``optuna.load_study()``, but for preferential optimization. + + Example: + + .. testsetup:: + + import os + + if os.path.exists("example.db"): + raise RuntimeError("'example.db' already exists. Please remove it.") + + .. testcode:: + + import optuna + from optuna_dashboard.preferential import create_study + from optuna_dashboard.preferential import load_study + + study = create_study(storage="sqlite:///example.db", study_name="my_study") + study.ask() + + loaded_study = load_study(study_name="my_study", storage="sqlite:///example.db") + assert len(loaded_study.trials) == len(study.trials) + + .. testcleanup:: + + os.remove("example.db") + + Args: + study_name: + Study's name. Each study has a unique name as an identifier. If :obj:`None`, checks + whether the storage contains a single study, and if so loads that study. + ``study_name`` is required if there are multiple studies in the storage. + storage: + Database URL such as ``sqlite:///example.db``. Please see also the documentation of + :func:`~optuna.study.create_study` for further details. + sampler: + A sampler object that implements background algorithm for value suggestion. + If :obj:`None` is specified, `RandomSampler`_ is used. Please note that + most Optuna samplers does not work efficiently for preferential optimization. + + .. _RandomSampler: https://optuna.readthedocs.io/en/stable/reference/samplers/generated/optuna.samplers.RandomSampler.html + + Returns: + A :class:`~optuna_dashboard.preferential.PreferentialStudy` object. + """ study = optuna.load_study( study_name=study_name, storage=storage, sampler=sampler or RandomSampler() ) From ae1d8f9b843f8679479ce75662c1dba02a06e203 Mon Sep 17 00:00:00 2001 From: c-bata Date: Mon, 14 Aug 2023 20:22:02 +0900 Subject: [PATCH 02/31] Fix lint errors --- optuna_dashboard/preferential/_study.py | 34 ++++++++++++++++--------- 1 file changed, 22 insertions(+), 12 deletions(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index 029a04e0..eca0303b 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -27,7 +27,8 @@ class PreferentialStudy: This object provides interfaces to create a new `Trial`_, set/get results of pairwise comparison called preferences. - .. _Trial: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.trial.Trial.html#optuna.trial.Trial + .. _Trial: https://optuna.readthedocs.io/en/stable/reference/generated/\ + optuna.trial.Trial.html#optuna.trial.Trial Note that the direct use of this constructor is not recommended. To create and load a study, please refer to the documentation of @@ -45,7 +46,8 @@ class PreferentialStudy: See `Study.trials`_ for details. - .. _Study.trials: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.trials + .. _Study.trials: https://optuna.readthedocs.io/en/stable/reference/generated/\ + optuna.study.Study.html#optuna.study.Study.trials Returns: A list of FrozenTrial object @@ -60,7 +62,8 @@ class PreferentialStudy: See `Study.best_trials`_ for details. - .. _Study.best_trials: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.best_trials + .. _Study.best_trials: https://optuna.readthedocs.io/en/stable/reference/\ + generated/optuna.study.Study.html#optuna.study.Study.best_trials Returns: A list of FrozenTrial object @@ -93,7 +96,8 @@ class PreferentialStudy: See `Study.user_attrs`_ for details. - .. _Study.user_attrs: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.user_attrs + .. _Study.user_attrs: https://optuna.readthedocs.io/en/stable/reference/\ + generated/optuna.study.Study.html#optuna.study.Study.user_attrs Returns: A dictionary containing all user attributes @@ -120,7 +124,8 @@ class PreferentialStudy: See `Study.get_trials`_ for details. - .. _Study.get_trials: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.get_trials + .. _Study.get_trials: https://optuna.readthedocs.io/en/stable/reference/\ + generated/optuna.study.Study.html#optuna.study.Study.get_trials Args: deepcopy: @@ -143,7 +148,8 @@ class PreferentialStudy: See `Study.ask`_ for details. - .. _Study.ask: https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.ask + .. _Study.ask: https://optuna.readthedocs.io/en/stable/reference/\ + generated/optuna.study.Study.html#optuna.study.Study.ask Args: fixed_distributions: @@ -164,7 +170,8 @@ class PreferentialStudy: See `Study.add_trials()`_ for details. - .. _Study.add_trials(): https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.add_trials + .. _Study.add_trials(): https://optuna.readthedocs.io/en/stable/reference/\ + generated/optuna.study.Study.html#optuna.study.Study.add_trials """ self._study.add_trial(trial) @@ -175,7 +182,8 @@ class PreferentialStudy: See `Study.add_trials()`_ for details. - .. _Study.add_trials(): https://optuna.readthedocs.io/en/stable/reference/generated/optuna.study.Study.html#optuna.study.Study.add_trials + .. _Study.add_trials(): https://optuna.readthedocs.io/en/stable/reference/\ + generated/optuna.study.Study.html#optuna.study.Study.add_trials """ self._study.add_trials(trials) @@ -223,8 +231,8 @@ class PreferentialStudy: .. seealso:: - See the `tutorial for user attributes `_on Optuna's documentation. + See the `tutorial for user attributes `_ on Optuna's documentation. """ self._study.set_user_attr(key, value) @@ -277,7 +285,8 @@ def create_study( If :obj:`None` is specified, `RandomSampler`_ is used. Please note that most Optuna samplers does not work efficiently for preferential optimization. - .. _RandomSampler: https://optuna.readthedocs.io/en/stable/reference/samplers/generated/optuna.samplers.RandomSampler.html + .. _RandomSampler: https://optuna.readthedocs.io/en/stable/reference/\ + samplers/generated/optuna.samplers.RandomSampler.html study_name: Study's name. If this argument is set to None, a unique name is generated @@ -368,7 +377,8 @@ def load_study( If :obj:`None` is specified, `RandomSampler`_ is used. Please note that most Optuna samplers does not work efficiently for preferential optimization. - .. _RandomSampler: https://optuna.readthedocs.io/en/stable/reference/samplers/generated/optuna.samplers.RandomSampler.html + .. _RandomSampler: https://optuna.readthedocs.io/en/stable/reference/samplers/\ + generated/optuna.samplers.RandomSampler.html Returns: A :class:`~optuna_dashboard.preferential.PreferentialStudy` object. From bc1b709475d68e5102fa161a2529a30ffd4d2e70 Mon Sep 17 00:00:00 2001 From: c-bata Date: Mon, 14 Aug 2023 20:31:44 +0900 Subject: [PATCH 03/31] Run black error --- optuna_dashboard/preferential/_study.py | 1 + 1 file changed, 1 insertion(+) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index eca0303b..944a4850 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -35,6 +35,7 @@ class PreferentialStudy: :func:`~optuna_dashboard.preferential.create_study` and :func:`~optuna_dashboard.preferential.load_study` respectively. """ + def __init__(self, study: optuna.Study) -> None: self._study = study From bab2c87ef71a2cbaf6622116e0133734a243a021 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Thu, 17 Aug 2023 14:12:27 +0900 Subject: [PATCH 04/31] Display TrialArtifact only when trial is running --- optuna_dashboard/ts/components/TrialList.tsx | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/optuna_dashboard/ts/components/TrialList.tsx b/optuna_dashboard/ts/components/TrialList.tsx index eb700c13..487f124f 100644 --- a/optuna_dashboard/ts/components/TrialList.tsx +++ b/optuna_dashboard/ts/components/TrialList.tsx @@ -153,6 +153,7 @@ const TrialListDetail: FC<{ const theme = useTheme() const action = actionCreator() const artifactEnabled = useRecoilValue(artifactIsAvailable) + const isRunningTrial = trial.state === "Running" || trial.state === "Waiting" const startMs = trial.datetime_start?.getTime() const completeMs = trial.datetime_complete?.getTime() @@ -319,7 +320,7 @@ const TrialListDetail: FC<{ value !== null ? renderInfo(key, value) : null )} - {artifactEnabled && } + {artifactEnabled && isRunningTrial && } ) } From b56db254f9dc74e68dce1de5426dd04b3807cadc Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Thu, 17 Aug 2023 14:33:07 +0900 Subject: [PATCH 05/31] Store metadeta in trial_system_attr --- optuna_dashboard/artifact/_backend.py | 34 ++++++++++++--------------- 1 file changed, 15 insertions(+), 19 deletions(-) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index eea1e988..f8f60569 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -102,8 +102,8 @@ def register_artifact_route( "mimetype": mimetype or DEFAULT_MIME_TYPE, "encoding": encoding, } - attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id - storage.set_study_system_attr(study_id, attr_key, json.dumps(artifact)) + attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id + storage.set_trial_system_attr(trial_id, attr_key, json.dumps(artifact)) response.status = 201 trial = storage.get_trial(trial_id) @@ -112,7 +112,7 @@ def register_artifact_route( return {"reason": "Invalid study_id or trial_id"} return { "artifact_id": artifact_id, - "artifacts": list_trial_artifacts(storage.get_study_system_attrs(study_id), trial), + "artifacts": list_trial_artifacts(storage.get_trial_system_attrs(trial_id), trial), } @app.delete("/api/artifacts///") @@ -123,8 +123,8 @@ def register_artifact_route( return {"reason": "Cannot access to the artifacts."} artifact_store.remove(artifact_id) - attr_key = _artifact_prefix(trial_id) + artifact_id - storage.set_study_system_attr(study_id, attr_key, json.dumps(None)) + attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id + storage.set_trial_system_attr(trial_id, attr_key, json.dumps(None)) response.status = 204 return {} @@ -178,24 +178,20 @@ def upload_artifact( "encoding": encoding or guess_encoding, "filename": filename, } - attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id - storage.set_study_system_attr(study_id, attr_key, json.dumps(artifact)) + attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id + storage.set_trial_system_attr(trial_id, attr_key, json.dumps(artifact)) with open(file_path, "rb") as f: backend.write(artifact_id, f) return artifact_id -def _artifact_prefix(trial_id: int) -> str: - return ARTIFACTS_ATTR_PREFIX + f"{trial_id}:" - - def get_artifact_meta( storage: BaseStorage, study_id: int, trial_id: int, artifact_id: str ) -> Optional[ArtifactMeta]: - study_system_attr = storage.get_study_system_attrs(study_id) - attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id - artifact_meta = study_system_attr.get(attr_key) + trial_system_attr = storage.get_trial_system_attrs(trial_id) + attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id + artifact_meta = trial_system_attr.get(attr_key) if artifact_meta is not None: return json.loads(artifact_meta) @@ -209,9 +205,9 @@ def get_artifact_meta( def delete_all_artifacts(backend: ArtifactStore, storage: BaseStorage, study_id: int) -> None: artifact_metas = [] - study_system_attrs = storage.get_study_system_attrs(study_id) for trial in storage.get_all_trials(study_id): - trial_artifacts = list_trial_artifacts(study_system_attrs, trial) + trial_system_attrs = storage.get_trial_system_attrs(trial._trial_id) + trial_artifacts = list_trial_artifacts(trial_system_attrs, trial) artifact_metas.extend(trial_artifacts) for meta in artifact_metas: @@ -219,12 +215,12 @@ def delete_all_artifacts(backend: ArtifactStore, storage: BaseStorage, study_id: def list_trial_artifacts( - study_system_attrs: dict[str, Any], trial: FrozenTrial + trial_system_attrs: dict[str, Any], trial: FrozenTrial ) -> list[ArtifactMeta]: dashboard_artifact_metas = [ json.loads(value) - for key, value in study_system_attrs.items() - if key.startswith(_artifact_prefix(trial._trial_id)) + for key, value in trial_system_attrs.items() + if key.startswith(ARTIFACTS_ATTR_PREFIX) ] # See https://github.com/optuna/optuna/blob/f827582a8/optuna/artifacts/_upload.py#L16 From 2cf12457d5d38eafc471117c9b754503621871fb Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Thu, 17 Aug 2023 15:56:11 +0900 Subject: [PATCH 06/31] Remove unused argument --- optuna_dashboard/artifact/_backend.py | 5 ++--- python_tests/artifact/test_optuna_compatibility.py | 1 - 2 files changed, 2 insertions(+), 4 deletions(-) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index f8f60569..a26e6af8 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -66,7 +66,7 @@ def register_artifact_route( if artifact_store is None: response.status = 400 # Bad Request return b"Cannot access to the artifacts." - artifact_dict = get_artifact_meta(storage, study_id, trial_id, artifact_id) + artifact_dict = get_artifact_meta(storage, trial_id, artifact_id) if artifact_dict is None: response.status = 404 return b"Not Found" @@ -169,7 +169,6 @@ def upload_artifact( filename = os.path.basename(file_path) storage = trial.storage trial_id = trial._trial_id - study_id = trial.study._study_id artifact_id = str(uuid.uuid4()) guess_mimetype, guess_encoding = mimetypes.guess_type(filename) artifact: ArtifactMeta = { @@ -187,7 +186,7 @@ def upload_artifact( def get_artifact_meta( - storage: BaseStorage, study_id: int, trial_id: int, artifact_id: str + storage: BaseStorage, trial_id: int, artifact_id: str ) -> Optional[ArtifactMeta]: trial_system_attr = storage.get_trial_system_attrs(trial_id) attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id diff --git a/python_tests/artifact/test_optuna_compatibility.py b/python_tests/artifact/test_optuna_compatibility.py index d81e1397..52032ae1 100644 --- a/python_tests/artifact/test_optuna_compatibility.py +++ b/python_tests/artifact/test_optuna_compatibility.py @@ -46,7 +46,6 @@ def test_list_optuna_trial_artifacts() -> None: artifact_meta = get_artifact_meta( storage=storage, - study_id=study._study_id, trial_id=trial._trial_id, artifact_id=artifact_id, ) From f00fd3b4e29e2d7773f973dd2c8fa23b794c2a51 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Thu, 17 Aug 2023 16:00:59 +0900 Subject: [PATCH 07/31] Check whether trial is finished or not when uploaded --- optuna_dashboard/artifact/_backend.py | 12 ++++++++---- 1 file changed, 8 insertions(+), 4 deletions(-) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index a26e6af8..b8ad54d8 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -81,6 +81,14 @@ def register_artifact_route( @app.post("/api/artifacts//") @json_api_view def upload_artifact_api(study_id: int, trial_id: int) -> dict[str, Any]: + trial = storage.get_trial(trial_id) + if trial is None: + response.status = 400 + return {"reason": "Invalid study_id or trial_id"} + elif trial.state.is_finished(): + response.status = 400 + return {"reason": "The trial is already finished."} + # TODO(c-bata): Use optuna.artifacts.upload_artifact() if artifact_store is None: response.status = 400 # Bad Request @@ -106,10 +114,6 @@ def register_artifact_route( storage.set_trial_system_attr(trial_id, attr_key, json.dumps(artifact)) response.status = 201 - trial = storage.get_trial(trial_id) - if trial is None: - response.status = 400 - return {"reason": "Invalid study_id or trial_id"} return { "artifact_id": artifact_id, "artifacts": list_trial_artifacts(storage.get_trial_system_attrs(trial_id), trial), From 10c1b76d155add8d9dea6ab1c4a285c54b375eb7 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Thu, 17 Aug 2023 17:05:46 +0900 Subject: [PATCH 08/31] Remove upload button only --- optuna_dashboard/ts/components/TrialList.tsx | 172 +++++++++---------- 1 file changed, 86 insertions(+), 86 deletions(-) diff --git a/optuna_dashboard/ts/components/TrialList.tsx b/optuna_dashboard/ts/components/TrialList.tsx index 487f124f..69c7de7a 100644 --- a/optuna_dashboard/ts/components/TrialList.tsx +++ b/optuna_dashboard/ts/components/TrialList.tsx @@ -1,3 +1,33 @@ +import CheckBoxIcon from "@mui/icons-material/CheckBox" +import CheckBoxOutlineBlankIcon from "@mui/icons-material/CheckBoxOutlineBlank" +import DeleteIcon from "@mui/icons-material/Delete" +import DownloadIcon from "@mui/icons-material/Download" +import FilterListIcon from "@mui/icons-material/FilterList" +import FullscreenIcon from "@mui/icons-material/Fullscreen" +import InsertDriveFileIcon from "@mui/icons-material/InsertDriveFile" +import StopCircleIcon from "@mui/icons-material/StopCircle" +import UploadFileIcon from "@mui/icons-material/UploadFile" +import { + Box, + Button, + Card, + CardActionArea, + CardContent, + CardMedia, + IconButton, + Menu, + MenuItem, + Modal, + Typography, + useTheme, +} from "@mui/material" +import Chip from "@mui/material/Chip" +import Divider from "@mui/material/Divider" +import List from "@mui/material/List" +import ListItem from "@mui/material/ListItem" +import ListItemButton from "@mui/material/ListItemButton" +import ListItemText from "@mui/material/ListItemText" +import ListSubheader from "@mui/material/ListSubheader" import React, { ChangeEventHandler, DragEventHandler, @@ -8,46 +38,16 @@ import React, { useRef, useState, } from "react" -import { - Typography, - Box, - Button, - useTheme, - IconButton, - Menu, - MenuItem, - Card, - CardContent, - CardMedia, - CardActionArea, - Modal, -} from "@mui/material" -import Chip from "@mui/material/Chip" -import Divider from "@mui/material/Divider" -import List from "@mui/material/List" -import ListItem from "@mui/material/ListItem" -import ListItemButton from "@mui/material/ListItemButton" -import ListItemText from "@mui/material/ListItemText" -import ListSubheader from "@mui/material/ListSubheader" -import FilterListIcon from "@mui/icons-material/FilterList" -import CheckBoxOutlineBlankIcon from "@mui/icons-material/CheckBoxOutlineBlank" -import CheckBoxIcon from "@mui/icons-material/CheckBox" -import UploadFileIcon from "@mui/icons-material/UploadFile" -import DownloadIcon from "@mui/icons-material/Download" -import DeleteIcon from "@mui/icons-material/Delete" -import FullscreenIcon from "@mui/icons-material/Fullscreen" -import InsertDriveFileIcon from "@mui/icons-material/InsertDriveFile" -import StopCircleIcon from "@mui/icons-material/StopCircle" -import { TrialNote } from "./Note" -import { useNavigate, useLocation } from "react-router-dom" import ListItemIcon from "@mui/material/ListItemIcon" +import { useLocation, useNavigate } from "react-router-dom" import { useRecoilValue } from "recoil" -import { artifactIsAvailable } from "../state" import { actionCreator } from "../action" +import { artifactIsAvailable } from "../state" import { useDeleteArtifactDialog } from "./DeleteArtifactDialog" -import { TrialFormWidgets } from "./TrialFormWidgets" +import { TrialNote } from "./Note" import { ThreejsArtifactViewer } from "./ThreejsArtifactViewer" +import { TrialFormWidgets } from "./TrialFormWidgets" const states: TrialState[] = [ "Complete", @@ -153,7 +153,6 @@ const TrialListDetail: FC<{ const theme = useTheme() const action = actionCreator() const artifactEnabled = useRecoilValue(artifactIsAvailable) - const isRunningTrial = trial.state === "Running" || trial.state === "Waiting" const startMs = trial.datetime_start?.getTime() const completeMs = trial.datetime_complete?.getTime() @@ -320,7 +319,7 @@ const TrialListDetail: FC<{ value !== null ? renderInfo(key, value) : null )} - {artifactEnabled && isRunningTrial && } + {artifactEnabled && } ) } @@ -711,55 +710,56 @@ const TrialArtifact: FC<{ trial: Trial }> = ({ trial }) => { ) } })} - - - - - - Upload a New File - - Drag your file here or click to browse. - - - - + + + Upload a New File + + Drag your file here or click to browse. + + + + + ) : null} {renderDeleteArtifactDialog()} @@ -952,15 +952,15 @@ export const TrialList: FC<{ studyDetail: StudyDetail | null }> = ({ {selected.length === 0 ? null : selected.map((t) => ( - - ))} + + ))} From 961f9c8716383124be95a33326d9472b96f6edf3 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Thu, 17 Aug 2023 20:33:17 +0900 Subject: [PATCH 09/31] Revert change on get_artifact_meta() --- python_tests/artifact/test_optuna_compatibility.py | 1 + 1 file changed, 1 insertion(+) diff --git a/python_tests/artifact/test_optuna_compatibility.py b/python_tests/artifact/test_optuna_compatibility.py index 52032ae1..d81e1397 100644 --- a/python_tests/artifact/test_optuna_compatibility.py +++ b/python_tests/artifact/test_optuna_compatibility.py @@ -46,6 +46,7 @@ def test_list_optuna_trial_artifacts() -> None: artifact_meta = get_artifact_meta( storage=storage, + study_id=study._study_id, trial_id=trial._trial_id, artifact_id=artifact_id, ) From 9e61d004aa418438e0df04c7414efea6a7977ec8 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Thu, 17 Aug 2023 23:00:50 +0900 Subject: [PATCH 10/31] Preserve backward compatibility when get and delete artifact --- optuna_dashboard/artifact/_backend.py | 47 ++++++++++++++++++--------- 1 file changed, 31 insertions(+), 16 deletions(-) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index b8ad54d8..b58c1daf 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -40,7 +40,7 @@ if TYPE_CHECKING: }, ) - +OPTUNA_ARTIFACTS_ATTR_PREFIX = "artifacts:" ARTIFACTS_ATTR_PREFIX = "dashboard:artifacts:" DEFAULT_MIME_TYPE = "application/octet-stream" BaseRequest.MEMFILE_MAX = int( @@ -66,7 +66,7 @@ def register_artifact_route( if artifact_store is None: response.status = 400 # Bad Request return b"Cannot access to the artifacts." - artifact_dict = get_artifact_meta(storage, trial_id, artifact_id) + artifact_dict = get_artifact_meta(storage, study_id, trial_id, artifact_id) if artifact_dict is None: response.status = 404 return b"Not Found" @@ -127,8 +127,11 @@ def register_artifact_route( return {"reason": "Cannot access to the artifacts."} artifact_store.remove(artifact_id) - attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id - storage.set_trial_system_attr(trial_id, attr_key, json.dumps(None)) + # The metadata of the artifact is stored in one of the following three locations: + storage.set_study_system_attr(study_id, _artifact_prefix(trial_id) + artifact_id, json.dumps(None)) + storage.set_trial_system_attr(trial_id, ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None)) + storage.set_trial_system_attr(trial_id, OPTUNA_ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None)) + response.status = 204 return {} @@ -189,28 +192,38 @@ def upload_artifact( return artifact_id +def _artifact_prefix(trial_id: int) -> str: + return ARTIFACTS_ATTR_PREFIX + f"{trial_id}:" + + def get_artifact_meta( - storage: BaseStorage, trial_id: int, artifact_id: str + storage: BaseStorage, study_id: int, trial_id: int, artifact_id: str ) -> Optional[ArtifactMeta]: - trial_system_attr = storage.get_trial_system_attrs(trial_id) - attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id - artifact_meta = trial_system_attr.get(attr_key) + # Search study_system_attrs due to backward compatibility. + study_system_attrs = storage.get_study_system_attrs(study_id) + attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id + artifact_meta = study_system_attrs.get(attr_key) if artifact_meta is not None: return json.loads(artifact_meta) + # Search trial_system_attrs. Note that artifacts uploaded via optuna.artifacts.upload_artifact + # have a different trial_system_attrs key prefix. # See https://github.com/optuna/optuna/blob/f827582a8/optuna/artifacts/_upload.py#L71 trial_system_attrs = storage.get_trial_system_attrs(trial_id) - value = trial_system_attrs.get("artifacts:" + artifact_id) + value = trial_system_attrs.get( + OPTUNA_ARTIFACTS_ATTR_PREFIX + artifact_id + ) or trial_system_attrs.get(ARTIFACTS_ATTR_PREFIX + artifact_id) if value is not None: return json.loads(value) + return None def delete_all_artifacts(backend: ArtifactStore, storage: BaseStorage, study_id: int) -> None: artifact_metas = [] + study_system_attrs = storage.get_study_system_attrs(study_id) for trial in storage.get_all_trials(study_id): - trial_system_attrs = storage.get_trial_system_attrs(trial._trial_id) - trial_artifacts = list_trial_artifacts(trial_system_attrs, trial) + trial_artifacts = list_trial_artifacts(study_system_attrs, trial) artifact_metas.extend(trial_artifacts) for meta in artifact_metas: @@ -218,20 +231,22 @@ def delete_all_artifacts(backend: ArtifactStore, storage: BaseStorage, study_id: def list_trial_artifacts( - trial_system_attrs: dict[str, Any], trial: FrozenTrial + study_system_attrs: dict[str, Any], trial: FrozenTrial ) -> list[ArtifactMeta]: + # Collect ArtifactMeta from study_system_attrs due to backward compatibility. dashboard_artifact_metas = [ json.loads(value) - for key, value in trial_system_attrs.items() - if key.startswith(ARTIFACTS_ATTR_PREFIX) + for key, value in study_system_attrs.items() + if key.startswith(_artifact_prefix(trial._trial_id)) ] + # Collect ArtifactMeta from trial_system_attrs. Note that artifacts uploaded via + # optuna.artifacts.upload_artifacts have a different trial_system_attrs key prefix. # See https://github.com/optuna/optuna/blob/f827582a8/optuna/artifacts/_upload.py#L16 optuna_artifact_metas = [ json.loads(value) for key, value in trial.system_attrs.items() - if key.startswith("artifacts:") + if key.startswith(OPTUNA_ARTIFACTS_ATTR_PREFIX) or key.startswith(ARTIFACTS_ATTR_PREFIX) ] - artifact_metas = dashboard_artifact_metas + optuna_artifact_metas return [a for a in artifact_metas if a is not None] From 5844697dbabc187df66809ee4485f1e09b8b9e17 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Fri, 18 Aug 2023 10:24:53 +0900 Subject: [PATCH 11/31] Apply black --- optuna_dashboard/artifact/_backend.py | 14 ++++++++++---- 1 file changed, 10 insertions(+), 4 deletions(-) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index b58c1daf..4ed07490 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -128,9 +128,15 @@ def register_artifact_route( artifact_store.remove(artifact_id) # The metadata of the artifact is stored in one of the following three locations: - storage.set_study_system_attr(study_id, _artifact_prefix(trial_id) + artifact_id, json.dumps(None)) - storage.set_trial_system_attr(trial_id, ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None)) - storage.set_trial_system_attr(trial_id, OPTUNA_ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None)) + storage.set_study_system_attr( + study_id, _artifact_prefix(trial_id) + artifact_id, json.dumps(None) + ) + storage.set_trial_system_attr( + trial_id, ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None) + ) + storage.set_trial_system_attr( + trial_id, OPTUNA_ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None) + ) response.status = 204 return {} @@ -215,7 +221,7 @@ def get_artifact_meta( ) or trial_system_attrs.get(ARTIFACTS_ATTR_PREFIX + artifact_id) if value is not None: return json.loads(value) - + return None From 571d1288b6cfcfeac45a388d09b59f1451d29268 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Fri, 18 Aug 2023 10:55:55 +0900 Subject: [PATCH 12/31] Apply fmt --- optuna_dashboard/ts/components/TrialList.tsx | 25 ++++++++++---------- 1 file changed, 13 insertions(+), 12 deletions(-) diff --git a/optuna_dashboard/ts/components/TrialList.tsx b/optuna_dashboard/ts/components/TrialList.tsx index 69c7de7a..f2ff1446 100644 --- a/optuna_dashboard/ts/components/TrialList.tsx +++ b/optuna_dashboard/ts/components/TrialList.tsx @@ -710,7 +710,7 @@ const TrialArtifact: FC<{ trial: Trial }> = ({ trial }) => { ) } })} - {(trial.state === "Running" || trial.state === "Waiting") ? ( + {trial.state === "Running" || trial.state === "Waiting" ? ( = ({ trial }) => { minHeight: height, margin: theme.spacing(0, 1, 1, 0), border: dragOver - ? `3px dashed ${theme.palette.mode === "dark" ? "white" : "black" - }` + ? `3px dashed ${ + theme.palette.mode === "dark" ? "white" : "black" + }` : `1px solid ${theme.palette.divider}`, }} onDragOver={handleDragOver} @@ -952,15 +953,15 @@ export const TrialList: FC<{ studyDetail: StudyDetail | null }> = ({ {selected.length === 0 ? null : selected.map((t) => ( - - ))} + + ))} From 9760436222d9951acdf595de3572cfb693e1e2f4 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Fri, 18 Aug 2023 10:59:24 +0900 Subject: [PATCH 13/31] Revert auto format --- optuna_dashboard/ts/components/TrialList.tsx | 74 ++++++++++---------- 1 file changed, 37 insertions(+), 37 deletions(-) diff --git a/optuna_dashboard/ts/components/TrialList.tsx b/optuna_dashboard/ts/components/TrialList.tsx index f2ff1446..4a54331d 100644 --- a/optuna_dashboard/ts/components/TrialList.tsx +++ b/optuna_dashboard/ts/components/TrialList.tsx @@ -1,33 +1,3 @@ -import CheckBoxIcon from "@mui/icons-material/CheckBox" -import CheckBoxOutlineBlankIcon from "@mui/icons-material/CheckBoxOutlineBlank" -import DeleteIcon from "@mui/icons-material/Delete" -import DownloadIcon from "@mui/icons-material/Download" -import FilterListIcon from "@mui/icons-material/FilterList" -import FullscreenIcon from "@mui/icons-material/Fullscreen" -import InsertDriveFileIcon from "@mui/icons-material/InsertDriveFile" -import StopCircleIcon from "@mui/icons-material/StopCircle" -import UploadFileIcon from "@mui/icons-material/UploadFile" -import { - Box, - Button, - Card, - CardActionArea, - CardContent, - CardMedia, - IconButton, - Menu, - MenuItem, - Modal, - Typography, - useTheme, -} from "@mui/material" -import Chip from "@mui/material/Chip" -import Divider from "@mui/material/Divider" -import List from "@mui/material/List" -import ListItem from "@mui/material/ListItem" -import ListItemButton from "@mui/material/ListItemButton" -import ListItemText from "@mui/material/ListItemText" -import ListSubheader from "@mui/material/ListSubheader" import React, { ChangeEventHandler, DragEventHandler, @@ -38,16 +8,46 @@ import React, { useRef, useState, } from "react" +import { + Typography, + Box, + Button, + useTheme, + IconButton, + Menu, + MenuItem, + Card, + CardContent, + CardMedia, + CardActionArea, + Modal, +} from "@mui/material" +import Chip from "@mui/material/Chip" +import Divider from "@mui/material/Divider" +import List from "@mui/material/List" +import ListItem from "@mui/material/ListItem" +import ListItemButton from "@mui/material/ListItemButton" +import ListItemText from "@mui/material/ListItemText" +import ListSubheader from "@mui/material/ListSubheader" +import FilterListIcon from "@mui/icons-material/FilterList" +import CheckBoxOutlineBlankIcon from "@mui/icons-material/CheckBoxOutlineBlank" +import CheckBoxIcon from "@mui/icons-material/CheckBox" +import UploadFileIcon from "@mui/icons-material/UploadFile" +import DownloadIcon from "@mui/icons-material/Download" +import DeleteIcon from "@mui/icons-material/Delete" +import FullscreenIcon from "@mui/icons-material/Fullscreen" +import InsertDriveFileIcon from "@mui/icons-material/InsertDriveFile" +import StopCircleIcon from "@mui/icons-material/StopCircle" -import ListItemIcon from "@mui/material/ListItemIcon" -import { useLocation, useNavigate } from "react-router-dom" -import { useRecoilValue } from "recoil" -import { actionCreator } from "../action" -import { artifactIsAvailable } from "../state" -import { useDeleteArtifactDialog } from "./DeleteArtifactDialog" import { TrialNote } from "./Note" -import { ThreejsArtifactViewer } from "./ThreejsArtifactViewer" +import { useNavigate, useLocation } from "react-router-dom" +import ListItemIcon from "@mui/material/ListItemIcon" +import { useRecoilValue } from "recoil" +import { artifactIsAvailable } from "../state" +import { actionCreator } from "../action" +import { useDeleteArtifactDialog } from "./DeleteArtifactDialog" import { TrialFormWidgets } from "./TrialFormWidgets" +import { ThreejsArtifactViewer } from "./ThreejsArtifactViewer" const states: TrialState[] = [ "Complete", From 408cf00d4a69f29d8a2b5c73f2f2b0b53ba05c8d Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Sat, 19 Aug 2023 10:55:57 +0900 Subject: [PATCH 14/31] Make ARTIFACTS_ATTR_PREFIX default consist with optuna --- optuna_dashboard/artifact/_backend.py | 14 +++++++------- 1 file changed, 7 insertions(+), 7 deletions(-) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index 4ed07490..5f0ba261 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -40,8 +40,8 @@ if TYPE_CHECKING: }, ) -OPTUNA_ARTIFACTS_ATTR_PREFIX = "artifacts:" -ARTIFACTS_ATTR_PREFIX = "dashboard:artifacts:" +ARTIFACTS_ATTR_PREFIX = "artifacts:" +DASHBOARD_ARTIFACTS_ATTR_PREFIX = "dashboard:artifacts:" DEFAULT_MIME_TYPE = "application/octet-stream" BaseRequest.MEMFILE_MAX = int( os.environ.get("OPTUNA_DASHBOARD_MEMFILE_MAX", 1024 * 1024 * 128) @@ -132,10 +132,10 @@ def register_artifact_route( study_id, _artifact_prefix(trial_id) + artifact_id, json.dumps(None) ) storage.set_trial_system_attr( - trial_id, ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None) + trial_id, DASHBOARD_ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None) ) storage.set_trial_system_attr( - trial_id, OPTUNA_ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None) + trial_id, ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None) ) response.status = 204 @@ -199,7 +199,7 @@ def upload_artifact( def _artifact_prefix(trial_id: int) -> str: - return ARTIFACTS_ATTR_PREFIX + f"{trial_id}:" + return DASHBOARD_ARTIFACTS_ATTR_PREFIX + f"{trial_id}:" def get_artifact_meta( @@ -217,7 +217,7 @@ def get_artifact_meta( # See https://github.com/optuna/optuna/blob/f827582a8/optuna/artifacts/_upload.py#L71 trial_system_attrs = storage.get_trial_system_attrs(trial_id) value = trial_system_attrs.get( - OPTUNA_ARTIFACTS_ATTR_PREFIX + artifact_id + DASHBOARD_ARTIFACTS_ATTR_PREFIX + artifact_id ) or trial_system_attrs.get(ARTIFACTS_ATTR_PREFIX + artifact_id) if value is not None: return json.loads(value) @@ -252,7 +252,7 @@ def list_trial_artifacts( optuna_artifact_metas = [ json.loads(value) for key, value in trial.system_attrs.items() - if key.startswith(OPTUNA_ARTIFACTS_ATTR_PREFIX) or key.startswith(ARTIFACTS_ATTR_PREFIX) + if key.startswith(ARTIFACTS_ATTR_PREFIX) or key.startswith(DASHBOARD_ARTIFACTS_ATTR_PREFIX) ] artifact_metas = dashboard_artifact_metas + optuna_artifact_metas return [a for a in artifact_metas if a is not None] From f6bc5ba77aa042be7e9826cd14f5e24c62e51421 Mon Sep 17 00:00:00 2001 From: keisuke-umezawa Date: Tue, 30 May 2023 22:36:30 +0900 Subject: [PATCH 15/31] Merge the implementation of GraphEdf and GraphEdfMultiStudies --- .../ts/components/CompareStudies.tsx | 32 ++-- optuna_dashboard/ts/components/GraphEdf.tsx | 164 ++---------------- .../ts/components/StudyDetail.tsx | 7 +- 3 files changed, 40 insertions(+), 163 deletions(-) diff --git a/optuna_dashboard/ts/components/CompareStudies.tsx b/optuna_dashboard/ts/components/CompareStudies.tsx index 99069205..7cb6bd85 100644 --- a/optuna_dashboard/ts/components/CompareStudies.tsx +++ b/optuna_dashboard/ts/components/CompareStudies.tsx @@ -12,6 +12,7 @@ import { useTheme, IconButton, } from "@mui/material" +import Grid2 from "@mui/material/Unstable_Grid2" import ChevronRightIcon from "@mui/icons-material/ChevronRight" import Chip from "@mui/material/Chip" import FormControlLabel from "@mui/material/FormControlLabel" @@ -325,19 +326,24 @@ const StudiesGraph: FC<{ studies: StudySummary[] }> = ({ studies }) => { ) : null} - {showStudyDetails !== null && - showStudyDetails.length > 0 && - showStudyDetails.every((s) => s) ? ( - - - - - - ) : null} + + {showStudyDetails !== null && + showStudyDetails.length > 0 && + showStudyDetails.every((s) => s) + ? showStudyDetails[0].directions.map((d, i) => ( + + + + + + + + )) + : null} + ) } diff --git a/optuna_dashboard/ts/components/GraphEdf.tsx b/optuna_dashboard/ts/components/GraphEdf.tsx index 185dd889..c0a1513c 100644 --- a/optuna_dashboard/ts/components/GraphEdf.tsx +++ b/optuna_dashboard/ts/components/GraphEdf.tsx @@ -1,25 +1,9 @@ import * as plotly from "plotly.js-dist-min" import React, { FC, useEffect, useMemo } from "react" -import { - Grid, - FormControl, - FormLabel, - MenuItem, - Select, - Typography, - SelectChangeEvent, - useTheme, - Box, -} from "@mui/material" +import { Typography, useTheme, Box } from "@mui/material" import { plotlyDarkTemplate } from "./PlotlyDarkMode" -import { - Target, - useFilteredTrials, - useFilteredTrialsFromStudies, - useObjectiveTargets, -} from "../trialFilter" +import { Target, useFilteredTrialsFromStudies } from "../trialFilter" -const plotDomId = "graph-edf" const getPlotDomId = (objectiveId: number) => `graph-edf-${objectiveId}` interface EdfPlotInfo { @@ -27,45 +11,17 @@ interface EdfPlotInfo { trials: Trial[] } -export const GraphEdf: FC<{ - study: StudyDetail | null +export const GraphEdfMultiStudies: FC<{ + studies: StudyDetail[] objectiveId: number -}> = ({ study, objectiveId }) => { +}> = ({ studies, objectiveId }) => { const theme = useTheme() const domId = getPlotDomId(objectiveId) const target = useMemo( () => new Target("objective", objectiveId), [objectiveId] ) - const trials = useFilteredTrials(study, [target], false) - - useEffect(() => { - if (study !== null) { - plotEdf(trials, target, domId, theme.palette.mode) - } - }, [trials, target, domId, theme.palette.mode]) - return ( - - - {`EDF for ${target.toLabel(study?.objective_names)}`} - - - - ) -} - -export const GraphEdfMultiStudies: FC<{ - studies: StudyDetail[] -}> = ({ studies }) => { - const theme = useTheme() - const [targets, selected, setTarget] = useObjectiveTargets( - studies.length !== 0 ? studies[0] : null - ) - - const trials = useFilteredTrialsFromStudies(studies, [selected], false) + const trials = useFilteredTrialsFromStudies(studies, [target], false) const edfPlotInfos = studies.map((study, index) => { const e: EdfPlotInfo = { study_name: study?.name, @@ -74,111 +30,23 @@ export const GraphEdfMultiStudies: FC<{ return e }) - const handleObjectiveChange = (event: SelectChangeEvent) => { - setTarget(event.target.value) - } - useEffect(() => { - plotEdfMultiStudies(edfPlotInfos, selected, plotDomId, theme.palette.mode) - }, [studies, selected, theme.palette.mode]) + plotEdfMultiStudies(edfPlotInfos, target, domId, theme.palette.mode) + }, [studies, target, theme.palette.mode]) return ( - - + - - EDF - - {studies.length > 0 && studies[0].directions.length !== 1 ? ( - - Objective: - - - ) : null} - - - - - + {`EDF for ${target.toLabel(studies[0].objective_names)}`} + + + ) } -const plotEdf = ( - trials: Trial[], - target: Target, - domId: string, - mode: string -) => { - if (document.getElementById(domId) === null) { - return - } - if (trials.length === 0) { - plotly.react(domId, [], { - template: mode === "dark" ? plotlyDarkTemplate : {}, - }) - return - } - - const target_name = "Objective Value" - const layout: Partial = { - xaxis: { - title: target_name, - }, - yaxis: { - title: "Cumulative Probability", - }, - margin: { - l: 50, - t: 0, - r: 50, - b: 50, - }, - uirevision: "true", - template: mode === "dark" ? plotlyDarkTemplate : {}, - } - - const values = trials.map((t) => target.getTargetValue(t) as number) - const numValues = values.length - const minX = Math.min(...values) - const maxX = Math.max(...values) - const numStep = 100 - const _step = (maxX - minX) / (numStep - 1) - - const xValues = [] - const yValues = [] - for (let i = 0; i < numStep; i++) { - const boundary_right = minX + _step * i - xValues.push(boundary_right) - yValues.push(values.filter((v) => v <= boundary_right).length / numValues) - } - - const plotData: Partial[] = [ - { - type: "scatter", - x: xValues, - y: yValues, - }, - ] - plotly.react(domId, plotData, layout) -} - const plotEdfMultiStudies = ( edfPlotInfos: EdfPlotInfo[], target: Target, diff --git a/optuna_dashboard/ts/components/StudyDetail.tsx b/optuna_dashboard/ts/components/StudyDetail.tsx index 255be7d2..9c167094 100644 --- a/optuna_dashboard/ts/components/StudyDetail.tsx +++ b/optuna_dashboard/ts/components/StudyDetail.tsx @@ -25,7 +25,7 @@ import { AppDrawer, PageId } from "./AppDrawer" import { GraphParallelCoordinate } from "./GraphParallelCoordinate" import { Contour } from "./GraphContour" import { GraphSlice } from "./GraphSlice" -import { GraphEdf } from "./GraphEdf" +import { GraphEdfMultiStudies } from "./GraphEdf" import { TrialList } from "./TrialList" import { StudyHistory } from "./StudyHistory" @@ -113,7 +113,10 @@ export const StudyDetail: FC<{ - + From 5fbbb16746b448340cd2feee415004969f305dcc Mon Sep 17 00:00:00 2001 From: keisuke-umezawa Date: Tue, 30 May 2023 22:40:55 +0900 Subject: [PATCH 16/31] Rename GraphEdfMultiStudies into GraphEdf --- optuna_dashboard/ts/components/CompareStudies.tsx | 9 +++------ optuna_dashboard/ts/components/GraphEdf.tsx | 6 +++--- optuna_dashboard/ts/components/StudyDetail.tsx | 7 ++----- 3 files changed, 8 insertions(+), 14 deletions(-) diff --git a/optuna_dashboard/ts/components/CompareStudies.tsx b/optuna_dashboard/ts/components/CompareStudies.tsx index 7cb6bd85..0354297e 100644 --- a/optuna_dashboard/ts/components/CompareStudies.tsx +++ b/optuna_dashboard/ts/components/CompareStudies.tsx @@ -27,8 +27,8 @@ import HomeIcon from "@mui/icons-material/Home" import { actionCreator } from "../action" import { studySummariesState, studyDetailsState } from "../state" import { AppDrawer } from "./AppDrawer" -import { GraphEdfMultiStudies } from "./GraphEdf" -import { GraphHistory } from "./GraphHistory" +import { GraphEdf } from "./GraphEdf" +import { GraphHistoryMultiStudies } from "./GraphHistory" import { useNavigate, useLocation } from "react-router-dom" const useQuery = (): URLSearchParams => { @@ -334,10 +334,7 @@ const StudiesGraph: FC<{ studies: StudySummary[] }> = ({ studies }) => { - + diff --git a/optuna_dashboard/ts/components/GraphEdf.tsx b/optuna_dashboard/ts/components/GraphEdf.tsx index c0a1513c..ca85ee54 100644 --- a/optuna_dashboard/ts/components/GraphEdf.tsx +++ b/optuna_dashboard/ts/components/GraphEdf.tsx @@ -11,7 +11,7 @@ interface EdfPlotInfo { trials: Trial[] } -export const GraphEdfMultiStudies: FC<{ +export const GraphEdf: FC<{ studies: StudyDetail[] objectiveId: number }> = ({ studies, objectiveId }) => { @@ -31,7 +31,7 @@ export const GraphEdfMultiStudies: FC<{ }) useEffect(() => { - plotEdfMultiStudies(edfPlotInfos, target, domId, theme.palette.mode) + plotEdf(edfPlotInfos, target, domId, theme.palette.mode) }, [studies, target, theme.palette.mode]) return ( @@ -47,7 +47,7 @@ export const GraphEdfMultiStudies: FC<{ ) } -const plotEdfMultiStudies = ( +const plotEdf = ( edfPlotInfos: EdfPlotInfo[], target: Target, domId: string, diff --git a/optuna_dashboard/ts/components/StudyDetail.tsx b/optuna_dashboard/ts/components/StudyDetail.tsx index 9c167094..10f32175 100644 --- a/optuna_dashboard/ts/components/StudyDetail.tsx +++ b/optuna_dashboard/ts/components/StudyDetail.tsx @@ -25,7 +25,7 @@ import { AppDrawer, PageId } from "./AppDrawer" import { GraphParallelCoordinate } from "./GraphParallelCoordinate" import { Contour } from "./GraphContour" import { GraphSlice } from "./GraphSlice" -import { GraphEdfMultiStudies } from "./GraphEdf" +import { GraphEdf } from "./GraphEdf" import { TrialList } from "./TrialList" import { StudyHistory } from "./StudyHistory" @@ -113,10 +113,7 @@ export const StudyDetail: FC<{ - + From 97eb633d7895117fe0fa202d285d9c8580abe33b Mon Sep 17 00:00:00 2001 From: keisuke-umezawa Date: Tue, 18 Jul 2023 16:20:30 +0900 Subject: [PATCH 17/31] Follow review comments --- .../ts/components/CompareStudies.tsx | 16 ++++++++-------- 1 file changed, 8 insertions(+), 8 deletions(-) diff --git a/optuna_dashboard/ts/components/CompareStudies.tsx b/optuna_dashboard/ts/components/CompareStudies.tsx index 0354297e..30dd9e0b 100644 --- a/optuna_dashboard/ts/components/CompareStudies.tsx +++ b/optuna_dashboard/ts/components/CompareStudies.tsx @@ -326,11 +326,11 @@ const StudiesGraph: FC<{ studies: StudySummary[] }> = ({ studies }) => { ) : null} - - {showStudyDetails !== null && - showStudyDetails.length > 0 && - showStudyDetails.every((s) => s) - ? showStudyDetails[0].directions.map((d, i) => ( + {showStudyDetails !== null && + showStudyDetails.length > 0 && + showStudyDetails.every((s) => s) + ? showStudyDetails[0].directions.map((d, i) => ( + @@ -338,9 +338,9 @@ const StudiesGraph: FC<{ studies: StudySummary[] }> = ({ studies }) => { - )) - : null} - + + )) + : null} ) } From 3ac858185cb05046cec9dc21bd808d8c5e713bac Mon Sep 17 00:00:00 2001 From: keisuke-umezawa Date: Sat, 19 Aug 2023 15:34:51 +0900 Subject: [PATCH 18/31] Fix type and console warnings --- .../ts/components/CompareStudies.tsx | 18 +++++++++--------- 1 file changed, 9 insertions(+), 9 deletions(-) diff --git a/optuna_dashboard/ts/components/CompareStudies.tsx b/optuna_dashboard/ts/components/CompareStudies.tsx index 30dd9e0b..e380485c 100644 --- a/optuna_dashboard/ts/components/CompareStudies.tsx +++ b/optuna_dashboard/ts/components/CompareStudies.tsx @@ -28,7 +28,7 @@ import { actionCreator } from "../action" import { studySummariesState, studyDetailsState } from "../state" import { AppDrawer } from "./AppDrawer" import { GraphEdf } from "./GraphEdf" -import { GraphHistoryMultiStudies } from "./GraphHistory" +import { GraphHistory } from "./GraphHistory" import { useNavigate, useLocation } from "react-router-dom" const useQuery = (): URLSearchParams => { @@ -326,11 +326,11 @@ const StudiesGraph: FC<{ studies: StudySummary[] }> = ({ studies }) => { ) : null} - {showStudyDetails !== null && - showStudyDetails.length > 0 && - showStudyDetails.every((s) => s) - ? showStudyDetails[0].directions.map((d, i) => ( - + + {showStudyDetails !== null && + showStudyDetails.length > 0 && + showStudyDetails.every((s) => s) + ? showStudyDetails[0].directions.map((d, i) => ( @@ -338,9 +338,9 @@ const StudiesGraph: FC<{ studies: StudySummary[] }> = ({ studies }) => { - - )) - : null} + )) + : null} + ) } From dae5f8f6e74363d5dddff395bb1a4b846f64e201 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Sat, 19 Aug 2023 16:42:15 +0900 Subject: [PATCH 19/31] Remove unnecessary search in trial_system_attr --- optuna_dashboard/artifact/_backend.py | 9 ++------- 1 file changed, 2 insertions(+), 7 deletions(-) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index 5f0ba261..7a9f1242 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -131,9 +131,6 @@ def register_artifact_route( storage.set_study_system_attr( study_id, _artifact_prefix(trial_id) + artifact_id, json.dumps(None) ) - storage.set_trial_system_attr( - trial_id, DASHBOARD_ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None) - ) storage.set_trial_system_attr( trial_id, ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None) ) @@ -216,9 +213,7 @@ def get_artifact_meta( # have a different trial_system_attrs key prefix. # See https://github.com/optuna/optuna/blob/f827582a8/optuna/artifacts/_upload.py#L71 trial_system_attrs = storage.get_trial_system_attrs(trial_id) - value = trial_system_attrs.get( - DASHBOARD_ARTIFACTS_ATTR_PREFIX + artifact_id - ) or trial_system_attrs.get(ARTIFACTS_ATTR_PREFIX + artifact_id) + value = trial_system_attrs.get(ARTIFACTS_ATTR_PREFIX + artifact_id) if value is not None: return json.loads(value) @@ -252,7 +247,7 @@ def list_trial_artifacts( optuna_artifact_metas = [ json.loads(value) for key, value in trial.system_attrs.items() - if key.startswith(ARTIFACTS_ATTR_PREFIX) or key.startswith(DASHBOARD_ARTIFACTS_ATTR_PREFIX) + if key.startswith(ARTIFACTS_ATTR_PREFIX) ] artifact_metas = dashboard_artifact_metas + optuna_artifact_metas return [a for a in artifact_metas if a is not None] From 98c85451242a2592bf20153b0c141b6776554ee3 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Sat, 19 Aug 2023 17:00:14 +0900 Subject: [PATCH 20/31] Update comment --- optuna_dashboard/artifact/_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index 7a9f1242..9ef79dad 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -127,7 +127,7 @@ def register_artifact_route( return {"reason": "Cannot access to the artifacts."} artifact_store.remove(artifact_id) - # The metadata of the artifact is stored in one of the following three locations: + # The artifact's metadata is stored in one of the following two locations: storage.set_study_system_attr( study_id, _artifact_prefix(trial_id) + artifact_id, json.dumps(None) ) From 182663158510741acb6bb0c5b660c28156a0a9ae Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Mon, 21 Aug 2023 12:32:45 +0900 Subject: [PATCH 21/31] Add unit tests --- python_tests/artifact/test_backend.py | 100 ++++++++++++++++++++++++++ 1 file changed, 100 insertions(+) create mode 100644 python_tests/artifact/test_backend.py diff --git a/python_tests/artifact/test_backend.py b/python_tests/artifact/test_backend.py new file mode 100644 index 00000000..d427d5d0 --- /dev/null +++ b/python_tests/artifact/test_backend.py @@ -0,0 +1,100 @@ +from unittest.mock import MagicMock + +from optuna_dashboard.artifact import _backend +import pytest + + +def test_get_artifact_path() -> None: + study = MagicMock(_study_id=0) + trial = MagicMock(_trial_id=0, study=study) + assert _backend.get_artifact_path(trial=trial, artifact_id="id0") == "/artifacts/0/0/id0" + + +def test_artifact_prefix() -> None: + actual = _backend._artifact_prefix(trial_id=0) + assert actual == "dashboard:artifacts:0:" + + +@pytest.fixture() +def init_storage_with_artifact_meta() -> MagicMock: + storage = MagicMock() + + get_study_system_attrs = ( + lambda study_id: { + "dashboard:artifacts:0:id0": '{"artifact_id": "id0", "filename": "foo.txt"}', + "dashboard:artifacts:0:id1": '{"artifact_id": "id1", "filename": "bar.txt"}', + "baz": "baz", + } + if study_id == 0 + else {} + ) + storage.get_study_system_attrs.side_effect = get_study_system_attrs + + trial0_system_attrs = { + "artifacts:id2": '{"artifact_id": "id2", "filename": "baz.txt"}', + } + trial1_system_attrs = { + "artifacts:id3": '{"artifact_id": "id3", "filename": "qux.txt"}', + } + get_trial_system_attrs = ( + lambda trial_id: trial0_system_attrs + if trial_id == 0 + else trial1_system_attrs + if trial_id == 1 + else {} + ) + storage.get_trial_system_attrs.side_effect = get_trial_system_attrs + + get_all_trials = ( + lambda study_id: [ + MagicMock( + _trial_id=0, study=MagicMock(_study_id=study_id), system_attrs=trial0_system_attrs + ), + MagicMock( + _trial_id=1, study=MagicMock(_study_id=study_id), system_attrs=trial1_system_attrs + ), + ] + if study_id == 0 + else [] + ) + storage.get_all_trials.side_effect = get_all_trials + + return storage + + +def test_get_artifact_meta(init_storage_with_artifact_meta: MagicMock) -> None: + storage = init_storage_with_artifact_meta + + actual = _backend.get_artifact_meta(storage, study_id=0, trial_id=0, artifact_id="id0") + assert actual == {"artifact_id": "id0", "filename": "foo.txt"} + + actual = _backend.get_artifact_meta(storage, study_id=0, trial_id=1, artifact_id="id3") + assert actual == {"artifact_id": "id3", "filename": "qux.txt"} + + actual = _backend.get_artifact_meta(storage, study_id=0, trial_id=0, artifact_id="id4") + assert actual is None + + +def test_delete_all_artifacts(init_storage_with_artifact_meta: MagicMock) -> None: + backend = MagicMock() + storage = init_storage_with_artifact_meta + _backend.delete_all_artifacts(backend, storage, study_id=0) + + assert backend.remove.call_args_list == [ + (("id0",),), + (("id1",),), + (("id2",),), + (("id3",),), + ] + + +def test_list_trial_artifacts(init_storage_with_artifact_meta: MagicMock) -> None: + storage = init_storage_with_artifact_meta + trial = MagicMock(_trial_id=0, system_attrs=storage.get_trial_system_attrs(0)) + + actual = _backend.list_trial_artifacts(storage.get_study_system_attrs(0), trial) + assert actual == [ + {"artifact_id": "id0", "filename": "foo.txt"}, + {"artifact_id": "id1", "filename": "bar.txt"}, + {"artifact_id": "id2", "filename": "baz.txt"}, + ] From 6322e76def3c6f326c0ab8fdeba08ed6931cd123 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Thu, 24 Aug 2023 13:05:43 +0900 Subject: [PATCH 22/31] Revert argument for list_trial_artifacts() in upload_artifact_api() Co-authored-by: c-bata --- optuna_dashboard/artifact/_backend.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index 9ef79dad..b83a8a0c 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -116,7 +116,7 @@ def register_artifact_route( return { "artifact_id": artifact_id, - "artifacts": list_trial_artifacts(storage.get_trial_system_attrs(trial_id), trial), + "artifacts": list_trial_artifacts(storage.get_study_system_attrs(study_id), trial), } @app.delete("/api/artifacts///") From 3b10ad1313853d998eb4edffe14f20fb1c976ad4 Mon Sep 17 00:00:00 2001 From: c-bata Date: Thu, 24 Aug 2023 15:18:27 +0900 Subject: [PATCH 23/31] Fix the bug while displaying study_user_attrs --- optuna_dashboard/_serializer.py | 1 + optuna_dashboard/ts/apiClient.ts | 2 ++ optuna_dashboard/ts/components/StudyHistory.tsx | 2 +- optuna_dashboard/ts/types/index.d.ts | 1 + 4 files changed, 5 insertions(+), 1 deletion(-) diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index a037290e..0fa04124 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -131,6 +131,7 @@ def serialize_study_detail( serialized: dict[str, Any] = { "name": summary.study_name, "directions": [d.name.lower() for d in summary.directions], + "user_attrs": serialize_attrs(summary.user_attrs), } system_attrs = getattr(summary, "system_attrs", {}) if summary.datetime_start is not None: diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index d9881715..5068dc3e 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -59,6 +59,7 @@ interface StudyDetailResponse { name: string datetime_start: string directions: StudyDirection[] + user_attrs: Attribute[] trials: TrialResponse[] best_trials: TrialResponse[] intersection_search_space: SearchSpaceItem[] @@ -93,6 +94,7 @@ export const getStudyDetailAPI = ( name: res.data.name, datetime_start: new Date(res.data.datetime_start), directions: res.data.directions, + user_attrs: res.data.user_attrs, trials: trials, best_trials: best_trials, union_search_space: res.data.union_search_space, diff --git a/optuna_dashboard/ts/components/StudyHistory.tsx b/optuna_dashboard/ts/components/StudyHistory.tsx index ce30fc44..b47c557a 100644 --- a/optuna_dashboard/ts/components/StudyHistory.tsx +++ b/optuna_dashboard/ts/components/StudyHistory.tsx @@ -39,7 +39,7 @@ export const StudyHistory: FC<{ studyId: number }> = ({ studyId }) => { setIncludePruned(!includePruned) } - const userAttrs = studySummary?.user_attrs || [] + const userAttrs = studySummary?.user_attrs || studyDetail?.user_attrs || [] const userAttrColumns: DataGridColumn[] = [ { field: "key", label: "Key", sortable: true }, { field: "value", label: "Value", sortable: true }, diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index 7f661525..1720cc6b 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -185,6 +185,7 @@ type StudyDetail = { id: number name: string directions: StudyDirection[] + user_attrs: Attribute[] datetime_start: Date best_trials: Trial[] trials: Trial[] From d11d777029d5cf3fd8a3db01dce84a0fbf64fa5f Mon Sep 17 00:00:00 2001 From: lucasmrdt Date: Thu, 24 Aug 2023 16:04:39 +0200 Subject: [PATCH 24/31] Fix the bug while renaming a study with maximize direction --- optuna_dashboard/_app.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 256ddbb6..8ee65701 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -139,8 +139,12 @@ def create_app( response.status = 404 # Not found return {"reason": f"study_id={study_id} is not found"} + summary = get_study_summary(storage, study_id) + if summary is None: + response.status = 500 + return {"reason": "Failed to load the study"} try: - dst_study = optuna.create_study(storage=storage, study_name=dst_study_name) + dst_study = optuna.create_study(storage=storage, study_name=dst_study_name, directions=summary.directions) dst_study.add_trials(src_study.get_trials(deepcopy=False)) except DuplicatedStudyError: response.status = 400 # Bad request From c6a10fc717d12ccc444c3e01b0573680ba9f2fa0 Mon Sep 17 00:00:00 2001 From: lucasmrdt Date: Thu, 24 Aug 2023 16:13:17 +0200 Subject: [PATCH 25/31] Fix lint error: E501 line too long --- optuna_dashboard/_app.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 8ee65701..6dab087a 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -144,7 +144,9 @@ def create_app( response.status = 500 return {"reason": "Failed to load the study"} try: - dst_study = optuna.create_study(storage=storage, study_name=dst_study_name, directions=summary.directions) + dst_study = optuna.create_study( + storage=storage, study_name=dst_study_name, directions=summary.directions + ) dst_study.add_trials(src_study.get_trials(deepcopy=False)) except DuplicatedStudyError: response.status = 400 # Bad request From af7abc913dd2f8a1dc97ce655f2bf6ee564c0352 Mon Sep 17 00:00:00 2001 From: Shinichi Hemmi <50256998+Alnusjaponica@users.noreply.github.com> Date: Fri, 25 Aug 2023 13:53:54 +0900 Subject: [PATCH 26/31] Replace MagicMock with actual storage --- python_tests/artifact/test_backend.py | 56 +++++++++------------------ 1 file changed, 19 insertions(+), 37 deletions(-) diff --git a/python_tests/artifact/test_backend.py b/python_tests/artifact/test_backend.py index d427d5d0..e2d0f152 100644 --- a/python_tests/artifact/test_backend.py +++ b/python_tests/artifact/test_backend.py @@ -1,5 +1,6 @@ from unittest.mock import MagicMock +from optuna.storages import BaseStorage from optuna_dashboard.artifact import _backend import pytest @@ -16,48 +17,29 @@ def test_artifact_prefix() -> None: @pytest.fixture() -def init_storage_with_artifact_meta() -> MagicMock: - storage = MagicMock() +def init_storage_with_artifact_meta() -> BaseStorage: + from optuna import create_study + from optuna.storages import InMemoryStorage - get_study_system_attrs = ( - lambda study_id: { - "dashboard:artifacts:0:id0": '{"artifact_id": "id0", "filename": "foo.txt"}', - "dashboard:artifacts:0:id1": '{"artifact_id": "id1", "filename": "bar.txt"}', - "baz": "baz", - } - if study_id == 0 - else {} - ) - storage.get_study_system_attrs.side_effect = get_study_system_attrs + storage = InMemoryStorage() + study = create_study(storage=storage) - trial0_system_attrs = { - "artifacts:id2": '{"artifact_id": "id2", "filename": "baz.txt"}', + study_system_attrs = { + "dashboard:artifacts:0:id0": '{"artifact_id": "id0", "filename": "foo.txt"}', + "dashboard:artifacts:0:id1": '{"artifact_id": "id1", "filename": "bar.txt"}', + "baz": "baz", } - trial1_system_attrs = { + for key, value in study_system_attrs.items(): + study.set_system_attr(key, value) + + trial_system_attrs = { + "artifacts:id2": '{"artifact_id": "id2", "filename": "baz.txt"}', "artifacts:id3": '{"artifact_id": "id3", "filename": "qux.txt"}', } - get_trial_system_attrs = ( - lambda trial_id: trial0_system_attrs - if trial_id == 0 - else trial1_system_attrs - if trial_id == 1 - else {} - ) - storage.get_trial_system_attrs.side_effect = get_trial_system_attrs - - get_all_trials = ( - lambda study_id: [ - MagicMock( - _trial_id=0, study=MagicMock(_study_id=study_id), system_attrs=trial0_system_attrs - ), - MagicMock( - _trial_id=1, study=MagicMock(_study_id=study_id), system_attrs=trial1_system_attrs - ), - ] - if study_id == 0 - else [] - ) - storage.get_all_trials.side_effect = get_all_trials + for key, value in trial_system_attrs.items(): + trial = study.ask() + trial.set_system_attr(key, value) + study.tell(trial, 0.0) return storage From 62dbdee9d0a29e4857dd3593d031088c2c9941da Mon Sep 17 00:00:00 2001 From: lucasmrdt Date: Fri, 25 Aug 2023 08:42:14 +0200 Subject: [PATCH 27/31] Get src directions from src_study instead of fetching study summary --- optuna_dashboard/_app.py | 6 +----- 1 file changed, 1 insertion(+), 5 deletions(-) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 6dab087a..2401f3f0 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -139,13 +139,9 @@ def create_app( response.status = 404 # Not found return {"reason": f"study_id={study_id} is not found"} - summary = get_study_summary(storage, study_id) - if summary is None: - response.status = 500 - return {"reason": "Failed to load the study"} try: dst_study = optuna.create_study( - storage=storage, study_name=dst_study_name, directions=summary.directions + storage=storage, study_name=dst_study_name, directions=src_study.directions ) dst_study.add_trials(src_study.get_trials(deepcopy=False)) except DuplicatedStudyError: From 7e68690a57bd2fcec160297680909b4c7b829f59 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 27 Aug 2023 01:18:12 +0900 Subject: [PATCH 28/31] Add some improvements on SQLite3 WASM loader --- standalone_app/src/sqlite3.ts | 327 ++++++++++++++++------------ standalone_app/src/types/index.d.ts | 12 +- 2 files changed, 191 insertions(+), 148 deletions(-) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index 4f65a65a..5996be5c 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -2,6 +2,14 @@ import sqlite3InitModule from "@sqlite.org/sqlite-wasm" import { SetterOrUpdater } from "recoil" +type SQLite3DB = { + exec(options: { + sql: string + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (...args: any[]) => void + }): void +} + export const loadStorage = ( arrayBuffer: ArrayBuffer, setter: SetterOrUpdater @@ -30,155 +38,186 @@ export const loadStorage = ( ) db.checkRc(rc) try { - // Check version_info table - let supported = true - db.exec({ - sql: "SELECT schema_version FROM version_info LIMIT 1", - // eslint-disable-next-line @typescript-eslint/no-explicit-any - callback: (vals: any[]) => { - if (vals[0] != 12) { - supported = false - } - }, - }) - if (!supported) { + if (!isSupportedSchema(db)) { return } - - // Get studies - const studies: Study[] = [] - db.exec({ - sql: - "SELECT s.study_id, s.study_name, sd.direction, sd.objective" + - " FROM studies AS s INNER JOIN study_directions AS sd" + - " ON s.study_id = sd.study_id ORDER BY sd.study_direction_id", - // eslint-disable-next-line @typescript-eslint/no-explicit-any - callback: (vals: any[]) => { - const study_id = vals[0] - const study_name = vals[1] - const direction: StudyDirection = vals[2].toLowerCase() - const objective = vals[3] - let index = 0 - - if (objective === 0) { - studies.push({ - study_id: study_id, - study_name: study_name, - directions: [direction], - union_search_space: [], - intersection_search_space: [], - user_attrs: [], - system_attrs: [], - trials: [], - }) - } else { - index = studies.findIndex((s) => s.study_id === study_id) - studies[index].directions.push(direction) - } - }, - }) - - studies.forEach((s) => { - db.exec({ - sql: - "SELECT t.trial_id, t.number, t.study_id, t.state, t.datetime_start, t.datetime_complete," + - " tv.objective, tv.value, tv.value_type" + - " FROM trials AS t LEFT JOIN trial_values AS tv ON tv.trial_id = t.trial_id" + - ` WHERE t.study_id = ${s.study_id}` + - " ORDER BY t.number", - // eslint-disable-next-line @typescript-eslint/no-explicit-any - callback: (vals: any[]) => { - const state: TrialState = - vals[3] === "COMPLETE" - ? "Complete" - : vals[3] === "PRUNED" - ? "Pruned" - : vals[3] === "RUNNING" - ? "Running" - : vals[3] === "WAITING" - ? "Waiting" - : "Fail" - const trial: Trial = { - trial_id: vals[0], - number: vals[1], - study_id: vals[2], - state: state, - params: [], - intermediate_values: [], - user_attrs: [], - system_attrs: [], - } - s.trials.push(trial) - }, - }) - const union_search_space: SearchSpaceItem[] = [] - let intersection_search_space: Set = new Set() - s.trials.forEach((trial) => { - const params: TrialParam[] = [] - const param_names = new Set() - db.exec({ - sql: - "SELECT param_name, param_value" + - ` FROM trial_params WHERE trial_id = ${trial.trial_id}`, - // eslint-disable-next-line @typescript-eslint/no-explicit-any - callback: (vals: any[]) => { - const param_name = vals[0] - params.push({ - name: param_name, - param_internal_value: vals[1], - }) - - param_names.add(param_name) - if ( - union_search_space.findIndex((s) => s.name === param_name) == -1 - ) { - union_search_space.push({ name: param_name }) - } - }, - }) - if (intersection_search_space.size === 0) { - param_names.forEach((s) => { - intersection_search_space.add({ - name: s, - }) - }) - } else { - intersection_search_space = new Set( - Array.from(intersection_search_space).filter((s) => - param_names.has(s.name) - ) - ) - } - - trial.params = params - const values: TrialValueNumber[] = [] - db.exec({ - sql: - "SELECT value, value_type" + - ` FROM trial_values WHERE trial_id = ${trial.trial_id}` + - " ORDER BY objective", - // eslint-disable-next-line @typescript-eslint/no-explicit-any - callback: (vals: any[]) => { - values.push( - vals[1] === "INF_NEG" - ? "-inf" - : vals[1] === "INF_POS" - ? "+inf" - : vals[0] - ) - }, - }) - if (s.directions.length === values.length) { - trial.values = values - } - }) - s.union_search_space = union_search_space - s.intersection_search_space = Array.from(intersection_search_space) - }) - + const studies = getStudies(db) setter((prev) => [...prev, ...studies]) } finally { db.close() } }) } + +const isSupportedSchema = (db: SQLite3DB): boolean => { + let supported = true + db.exec({ + sql: "SELECT schema_version FROM version_info LIMIT 1", + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + if (vals[0] != 12) { + supported = false + } + }, + }) + return supported +} + +const getStudies = (db: SQLite3DB): Study[] => { + const studies: Study[] = [] + db.exec({ + sql: + "SELECT s.study_id, s.study_name, sd.direction, sd.objective" + + " FROM studies AS s INNER JOIN study_directions AS sd" + + " ON s.study_id = sd.study_id ORDER BY sd.study_direction_id", + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + const studyId = vals[0] + const studyName = vals[1] + const direction: StudyDirection = + vals[2] === "MINIMIZE" ? "minimize" : "maximize" + const objective = vals[3] + + const trials = getTrials(db, studyId) + const union_search_space: SearchSpaceItem[] = [] + let intersection_search_space: Set = new Set() + trials.forEach((trial) => { + const params: TrialParam[] = [] + const param_names = new Set() + db.exec({ + sql: + "SELECT param_name, param_value" + + ` FROM trial_params WHERE trial_id = ${trial.trial_id}`, + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + const param_name = vals[0] + // TODO(c-bata): Support param_external_value + params.push({ + name: param_name, + param_internal_value: vals[1], + }) + + param_names.add(param_name) + if ( + union_search_space.findIndex((s) => s.name === param_name) == -1 + ) { + union_search_space.push({ name: param_name }) + } + }, + }) + if (intersection_search_space.size === 0) { + param_names.forEach((s) => { + intersection_search_space.add({ + name: s, + }) + }) + } else { + intersection_search_space = new Set( + Array.from(intersection_search_space).filter((s) => + param_names.has(s.name) + ) + ) + } + trial.params = params + }) + + if (objective === 0) { + studies.push({ + study_id: studyId, + study_name: studyName, + directions: [direction], + union_search_space: union_search_space, + intersection_search_space: Array.from(intersection_search_space), + trials: trials, + }) + return + } + const index = studies.findIndex((s) => s.study_id === studyId) + studies[index].directions.push(direction) + }, + }) + return studies +} + +const getTrials = (db: SQLite3DB, studyId: number): Trial[] => { + const trials: Trial[] = [] + db.exec({ + sql: + "SELECT trial_id, number, state, datetime_start, datetime_complete FROM trials" + + ` WHERE study_id = ${studyId} ORDER BY number`, + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + const trialId = vals[0] + const state: TrialState = + vals[2] === "COMPLETE" + ? "Complete" + : vals[2] === "PRUNED" + ? "Pruned" + : vals[2] === "RUNNING" + ? "Running" + : vals[2] === "WAITING" + ? "Waiting" + : "Fail" + const trial: Trial = { + trial_id: trialId, + number: vals[1], + study_id: studyId, + state: state, + values: getTrialValues(db, trialId), + params: [], // Set this column later + intermediate_values: getTrialIntermediateValues(db, trialId), + datetime_start: vals[3], + datetime_complete: vals[4], + } + trials.push(trial) + }, + }) + return trials +} + +const getTrialValues = (db: SQLite3DB, trialId: number): TrialValueNumber[] => { + const values: TrialValueNumber[] = [] + db.exec({ + sql: + "SELECT value, value_type" + + ` FROM trial_values WHERE trial_id = ${trialId}` + + " ORDER BY objective", + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + values.push( + vals[1] === "INF_NEG" + ? "-inf" + : vals[1] === "INF_POS" + ? "+inf" + : vals[0] + ) + }, + }) + return values +} + +const getTrialIntermediateValues = ( + db: SQLite3DB, + trialId: number +): TrialIntermediateValue[] => { + const values: TrialIntermediateValue[] = [] + db.exec({ + sql: + "SELECT step, intermediate_value, intermediate_value_type" + + ` FROM trial_values WHERE trial_id = ${trialId}` + + " ORDER BY step", + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + values.push({ + step: vals[0], + value: + vals[2] === "INF_NEG" + ? "-inf" + : vals[2] === "INF_POS" + ? "+inf" + : vals[1], + }) + }, + }) + return values +} diff --git a/standalone_app/src/types/index.d.ts b/standalone_app/src/types/index.d.ts index 0b676221..a0f22bbb 100644 --- a/standalone_app/src/types/index.d.ts +++ b/standalone_app/src/types/index.d.ts @@ -27,6 +27,11 @@ type CategoricalDistribution = { choices: { pytype: string; value: string }[] } +type TrialIntermediateValue = { + step: number + value: TrialIntermediateValueNumber +} + type Distribution = | FloatDistribution | IntDistribution @@ -41,10 +46,8 @@ type Study = { study_id: number study_name: string directions: StudyDirection[] - user_attrs: Attribute[] union_search_space: SearchSpaceItem[] intersection_search_space: SearchSpaceItem[] - system_attrs: Attribute[] datetime_start?: Date trials: Trial[] } @@ -59,13 +62,14 @@ type Trial = { intermediate_values: TrialIntermediateValue[] datetime_start?: Date datetime_complete?: Date - user_attrs: Attribute[] - system_attrs: Attribute[] } type TrialParam = { name: string param_internal_value: number + // param_external_value: string + // param_external_type: string + // distribution: Distribution } type SearchSpaceItem = { From 223f91dfbdc4efc1164cfe7582037a961917fd80 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 27 Aug 2023 02:50:01 +0900 Subject: [PATCH 29/31] Support params and union_user_attrs --- standalone_app/src/components/TrialTable.tsx | 90 ++++++++----- standalone_app/src/sqlite3.ts | 134 +++++++++++++++---- standalone_app/src/types/index.d.ts | 13 +- 3 files changed, 181 insertions(+), 56 deletions(-) diff --git a/standalone_app/src/components/TrialTable.tsx b/standalone_app/src/components/TrialTable.tsx index 97c14ea0..58d2d101 100644 --- a/standalone_app/src/components/TrialTable.tsx +++ b/standalone_app/src/components/TrialTable.tsx @@ -20,36 +20,6 @@ export const TrialTable: FC<{ }, ] - study.union_search_space.forEach((s) => { - columns.push({ - field: "params", - label: `Param ${s.name}`, - toCellValue: (i) => - trials[i].params.find((p) => p.name === s.name)?.param_internal_value || - null, - sortable: true, - filterable: false, - less: (firstEl, secondEl): number => { - const firstVal = firstEl.params.find( - (p) => p.name === s.name - )?.param_internal_value - const secondVal = secondEl.params.find( - (p) => p.name === s.name - )?.param_internal_value - - if (firstVal === secondVal) { - return 0 - } else if (firstVal && secondVal) { - return firstVal < secondVal ? 1 : -1 - } else if (firstVal) { - return -1 - } else { - return 1 - } - }, - }) - }) - if (study === null || study.directions.length == 1) { columns.push({ field: "values", @@ -117,6 +87,66 @@ export const TrialTable: FC<{ columns.push(...objectiveColumns) } + study.union_search_space.forEach((s) => { + columns.push({ + field: "params", + label: `Param ${s.name}`, + toCellValue: (i) => + trials[i].params.find((p) => p.name === s.name)?.param_external_value ?? + null, + sortable: true, + filterable: false, + less: (firstEl, secondEl): number => { + const firstVal = firstEl.params.find( + (p) => p.name === s.name + )?.param_internal_value + const secondVal = secondEl.params.find( + (p) => p.name === s.name + )?.param_internal_value + + if (firstVal === secondVal) { + return 0 + } else if (firstVal && secondVal) { + return firstVal < secondVal ? 1 : -1 + } else if (firstVal) { + return -1 + } else { + return 1 + } + }, + }) + }) + + study.union_user_attrs.forEach((attr_spec) => { + columns.push({ + field: "user_attrs", + label: `UserAttribute ${attr_spec.key}`, + toCellValue: (i) => + trials[i].user_attrs.find((attr) => attr.key === attr_spec.key) + ?.value || null, + sortable: attr_spec.sortable, + filterable: false, + less: (firstEl, secondEl): number => { + const firstVal = firstEl.user_attrs.find( + (attr) => attr.key === attr_spec.key + )?.value + const secondVal = secondEl.user_attrs.find( + (attr) => attr.key === attr_spec.key + )?.value + + if (firstVal === secondVal) { + return 0 + } else if (firstVal && secondVal) { + return firstVal < secondVal ? 1 : -1 + } else if (firstVal) { + return -1 + } else { + return 1 + } + }, + }) + }) + return ( columns={columns} diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index 5996be5c..680fe9f3 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -80,30 +80,25 @@ const getStudies = (db: SQLite3DB): Study[] => { const trials = getTrials(db, studyId) const union_search_space: SearchSpaceItem[] = [] + const union_user_attrs: AttributeSpec[] = [] let intersection_search_space: Set = new Set() trials.forEach((trial) => { - const params: TrialParam[] = [] - const param_names = new Set() - db.exec({ - sql: - "SELECT param_name, param_value" + - ` FROM trial_params WHERE trial_id = ${trial.trial_id}`, - // eslint-disable-next-line @typescript-eslint/no-explicit-any - callback: (vals: any[]) => { - const param_name = vals[0] - // TODO(c-bata): Support param_external_value - params.push({ - name: param_name, - param_internal_value: vals[1], - }) + const userAttrs = getTrialUserAttributes(db, trial.trial_id) + userAttrs.forEach((attr) => { + if (union_user_attrs.findIndex((s) => s.key === attr.key) == -1) { + union_user_attrs.push({ key: attr.key, sortable: false }) + } + }) - param_names.add(param_name) - if ( - union_search_space.findIndex((s) => s.name === param_name) == -1 - ) { - union_search_space.push({ name: param_name }) - } - }, + const params = getTrialParams(db, trial.trial_id) + const param_names = new Set() + params.forEach((param) => { + param_names.add(param.name) + if ( + union_search_space.findIndex((s) => s.name === param.name) == -1 + ) { + union_search_space.push({ name: param.name }) + } }) if (intersection_search_space.size === 0) { param_names.forEach((s) => { @@ -119,6 +114,7 @@ const getStudies = (db: SQLite3DB): Study[] => { ) } trial.params = params + trial.user_attrs = userAttrs }) if (objective === 0) { @@ -128,6 +124,7 @@ const getStudies = (db: SQLite3DB): Study[] => { directions: [direction], union_search_space: union_search_space, intersection_search_space: Array.from(intersection_search_space), + union_user_attrs: union_user_attrs, trials: trials, }) return @@ -164,8 +161,9 @@ const getTrials = (db: SQLite3DB, studyId: number): Trial[] => { study_id: studyId, state: state, values: getTrialValues(db, trialId), - params: [], // Set this column later intermediate_values: getTrialIntermediateValues(db, trialId), + params: [], // Set this column later + user_attrs: [], // Set this column later datetime_start: vals[3], datetime_complete: vals[4], } @@ -196,6 +194,96 @@ const getTrialValues = (db: SQLite3DB, trialId: number): TrialValueNumber[] => { return values } +const getTrialParams = (db: SQLite3DB, trialId: number): TrialParam[] => { + const params: TrialParam[] = [] + db.exec({ + sql: + "SELECT param_name, param_value, distribution_json" + + ` FROM trial_params WHERE trial_id = ${trialId}`, + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + const distribution = parseDistributionJSON(vals[2]) + params.push({ + name: vals[0], + param_internal_value: vals[1], + param_external_type: distribution.type, + param_external_value: paramInternalValueToExternalValue( + distribution, + vals[1] + ), + distribution: distribution, + }) + }, + }) + return params +} + +const paramInternalValueToExternalValue = ( + distribution: Distribution, + internalValue: number +): string => { + if (distribution.type === "FloatDistribution") { + return internalValue.toString() + } else if (distribution.type === "IntDistribution") { + return internalValue.toString() + } else { + return distribution.choices[internalValue].value + } +} + +const parseDistributionJSON = (t: string): Distribution => { + const parsed = JSON.parse(t) + if (parsed.name === "FloatDistribution") { + return { + type: "FloatDistribution", + low: parsed.attributes.low as number, + high: parsed.attributes.high as number, + step: parsed.attributes.step as number, + log: parsed.attributes.log as boolean, + } + } else if (parsed.name === "IntDistribution") { + return { + type: "IntDistribution", + low: parsed.attributes.low as number, + high: parsed.attributes.high as number, + step: parsed.attributes.step as number, + log: parsed.attributes.log as boolean, + } + } else { + const choices = parsed.attributes.choices.map((value: any) => { + // TODO(c-bata): Support other types + return { + pytype: "str", + value: value.toString(), + } + }) + return { + type: "CategoricalDistribution", + choices: choices, + } + } +} + +const getTrialUserAttributes = ( + db: SQLite3DB, + trialId: number +): Attribute[] => { + const attrs: Attribute[] = [] + db.exec({ + sql: + "SELECT key, value_json" + + ` FROM trial_user_attributes WHERE trial_id = ${trialId}`, + // eslint-disable-next-line @typescript-eslint/no-explicit-any + callback: (vals: any[]) => { + attrs.push({ + key: vals[0], + value: vals[1], + }) + }, + }) + return attrs +} + const getTrialIntermediateValues = ( db: SQLite3DB, trialId: number @@ -204,7 +292,7 @@ const getTrialIntermediateValues = ( db.exec({ sql: "SELECT step, intermediate_value, intermediate_value_type" + - ` FROM trial_values WHERE trial_id = ${trialId}` + + ` FROM trial_intermediate_values WHERE trial_id = ${trialId}` + " ORDER BY step", // eslint-disable-next-line @typescript-eslint/no-explicit-any callback: (vals: any[]) => { diff --git a/standalone_app/src/types/index.d.ts b/standalone_app/src/types/index.d.ts index a0f22bbb..14daaaca 100644 --- a/standalone_app/src/types/index.d.ts +++ b/standalone_app/src/types/index.d.ts @@ -42,12 +42,18 @@ type Attribute = { value: string } +type AttributeSpec = { + key: string + sortable: boolean +} + type Study = { study_id: number study_name: string directions: StudyDirection[] union_search_space: SearchSpaceItem[] intersection_search_space: SearchSpaceItem[] + union_user_attrs: AttributeSpec[] datetime_start?: Date trials: Trial[] } @@ -60,6 +66,7 @@ type Trial = { values?: TrialValueNumber[] params: TrialParam[] intermediate_values: TrialIntermediateValue[] + user_attrs: Attribute[] datetime_start?: Date datetime_complete?: Date } @@ -67,9 +74,9 @@ type Trial = { type TrialParam = { name: string param_internal_value: number - // param_external_value: string - // param_external_type: string - // distribution: Distribution + param_external_value: string + param_external_type: string + distribution: Distribution } type SearchSpaceItem = { From 4fc3259b1c62bde4b5cfac19f3a3dfd2cc5c75b0 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 27 Aug 2023 02:55:46 +0900 Subject: [PATCH 30/31] Fix lint errors --- standalone_app/src/sqlite3.ts | 1 + 1 file changed, 1 insertion(+) diff --git a/standalone_app/src/sqlite3.ts b/standalone_app/src/sqlite3.ts index 680fe9f3..00570f9a 100644 --- a/standalone_app/src/sqlite3.ts +++ b/standalone_app/src/sqlite3.ts @@ -250,6 +250,7 @@ const parseDistributionJSON = (t: string): Distribution => { log: parsed.attributes.log as boolean, } } else { + // eslint-disable-next-line @typescript-eslint/no-explicit-any const choices = parsed.attributes.choices.map((value: any) => { // TODO(c-bata): Support other types return { From 98db8c34463ad44a550902167e3206fe5c682d42 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 27 Aug 2023 03:08:00 +0900 Subject: [PATCH 31/31] Add Intermediate values plot --- .../src/components/PlotIntermediateValues.tsx | 100 ++++++++++++++++++ standalone_app/src/components/StudyDetail.tsx | 33 ++++-- 2 files changed, 126 insertions(+), 7 deletions(-) create mode 100644 standalone_app/src/components/PlotIntermediateValues.tsx diff --git a/standalone_app/src/components/PlotIntermediateValues.tsx b/standalone_app/src/components/PlotIntermediateValues.tsx new file mode 100644 index 00000000..4feb0b33 --- /dev/null +++ b/standalone_app/src/components/PlotIntermediateValues.tsx @@ -0,0 +1,100 @@ +import * as plotly from "plotly.js-dist-min" +import React, { FC, useEffect } from "react" +import { Box, Typography, useTheme, CardContent, Card } from "@mui/material" +import { plotlyDarkTemplate } from "../PlotlyDarkMode" + +const plotDomId = "graph-intermediate-values" + +export const PlotIntermediateValues: FC<{ + trials: Trial[] + includePruned: boolean + logScale: boolean +}> = ({ trials, includePruned, logScale }) => { + const theme = useTheme() + + useEffect(() => { + plotIntermediateValue( + trials, + theme.palette.mode, + false, + !includePruned, + logScale + ) + }, [trials, theme.palette.mode, false, includePruned, logScale]) + + return ( + + + + Intermediate values + + + + + ) +} + +const plotIntermediateValue = ( + trials: Trial[], + mode: string, + filterCompleteTrial: boolean, + filterPrunedTrial: boolean, + logScale: boolean +) => { + if (document.getElementById(plotDomId) === null) { + return + } + + const layout: Partial = { + margin: { + l: 50, + t: 0, + r: 50, + b: 0, + }, + yaxis: { + title: "Objective Value", + type: logScale ? "log" : "linear", + }, + xaxis: { + title: "Step", + type: "linear", + }, + uirevision: "true", + template: mode === "dark" ? plotlyDarkTemplate : {}, + } + if (trials.length === 0) { + plotly.react(plotDomId, [], layout) + return + } + + const filteredTrials = trials.filter( + (t) => + (!filterCompleteTrial && t.state === "Complete") || + (!filterPrunedTrial && + t.state === "Pruned" && + t.values && + t.values.length > 0) || + t.state == "Running" + ) + const plotData: Partial[] = filteredTrials.map((trial) => { + const values = trial.intermediate_values.filter( + (iv) => iv.value !== "inf" && iv.value !== "-inf" && iv.value !== "nan" + ) + return { + x: values.map((iv) => iv.step), + y: values.map((iv) => iv.value), + marker: { maxdisplayed: 10 }, + mode: "lines+markers", + type: "scatter", + name: + trial.state !== "Running" + ? `trial #${trial.number}` + : `trial #${trial.number} (running)`, + } + }) + plotly.react(plotDomId, plotData, layout) +} diff --git a/standalone_app/src/components/StudyDetail.tsx b/standalone_app/src/components/StudyDetail.tsx index 5ba84239..c1619cb5 100644 --- a/standalone_app/src/components/StudyDetail.tsx +++ b/standalone_app/src/components/StudyDetail.tsx @@ -11,6 +11,7 @@ import { Card, CardContent, } from "@mui/material" +import Grid2 from "@mui/material/Unstable_Grid2" import { Home } from "@mui/icons-material" import Brightness4Icon from "@mui/icons-material/Brightness4" import Brightness7Icon from "@mui/icons-material/Brightness7" @@ -19,6 +20,7 @@ import { studiesState } from "../state" import { TrialTable } from "./TrialTable" import { PlotHistory } from "./PlotHistory" import { PlotImportance } from "./PlotImportance" +import { PlotIntermediateValues } from "./PlotIntermediateValues" const useStudyValue = (idx: number): Study | null => { const studies = useRecoilValue(studiesState) @@ -83,7 +85,7 @@ export const StudyDetail: FC<{ }, }} > -
+ <> - - - {!!study && } - - + + + + + {!!study && } + + + + + + + {!!study && ( + + )} + + + + {!!study && } -
+ )