From 549525f81da4ffa93947c41d4a2f077dc7915b56 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Thu, 31 Aug 2023 18:28:19 +0900 Subject: [PATCH 01/26] Separate best_trials and active_trials --- optuna_dashboard/preferential/_study.py | 51 +++++++++++++++++-------- 1 file changed, 35 insertions(+), 16 deletions(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index 16c81b66..5ea5bf36 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -62,18 +62,20 @@ class PreferentialStudy: 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 """ return get_best_trials(self._study._study_id, self._study._storage) + @property + def active_trials(self) -> list[FrozenTrial]: + """Return the trials that is not reported bad and not marked skipped. + + Returns: + A list of FrozenTrial object + """ + return get_active_trials(self._study._study_id, self._study._storage) + @property def study_name(self) -> str: """Return the name of the study. @@ -283,21 +285,38 @@ class PreferentialStudy: to generate a new trial if this method returns :obj:`True`, and to wait for human evaluation if this method returns :obj:`False`. """ - return len(self.best_trials) < self.n_generate + return len(self.active_trials) < self.n_generate + + +def get_active_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]: + preferences = get_preferences(study_id, storage) + worse_numbers = {worse for _, worse in preferences} + study_system_attrs = storage.get_study_system_attrs(study_id) + active_trials = [] + for t in storage.get_all_trials( + study_id, deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING) + ): + if t.number in worse_numbers: + continue + if is_skipped_trial(t._trial_id, study_system_attrs): + continue + active_trials.append(copy.deepcopy(t)) + return active_trials def get_best_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]: preferences = get_preferences(study_id, storage) worse_numbers = {worse for _, worse in preferences} - study_system_attrs = storage.get_study_system_attrs(study_id) - best_trials = [] - for t in storage.get_all_trials( + nondominated_numbers = {better for better, _ in preferences if better not in worse_numbers} + trials = storage.get_all_trials( study_id, deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING) - ): - if not t.system_attrs.get(_SYSTEM_ATTR_COMPARISON_READY, False): - continue - if t.number in worse_numbers: - continue + ) + + study_system_attrs = storage.get_study_system_attrs(study_id) + + best_trials = [] + for n in nondominated_numbers: + t = trials[n] if is_skipped_trial(t._trial_id, study_system_attrs): continue best_trials.append(copy.deepcopy(t)) From 6553bb5fbdeb25d92ee99099092e589b3596eea5 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Thu, 31 Aug 2023 18:37:46 +0900 Subject: [PATCH 02/26] Remove PreferentialStudy.active_trial --- optuna_dashboard/preferential/_study.py | 11 +---------- 1 file changed, 1 insertion(+), 10 deletions(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index 5ea5bf36..57ab8e9f 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -67,15 +67,6 @@ class PreferentialStudy: """ return get_best_trials(self._study._study_id, self._study._storage) - @property - def active_trials(self) -> list[FrozenTrial]: - """Return the trials that is not reported bad and not marked skipped. - - Returns: - A list of FrozenTrial object - """ - return get_active_trials(self._study._study_id, self._study._storage) - @property def study_name(self) -> str: """Return the name of the study. @@ -285,7 +276,7 @@ class PreferentialStudy: to generate a new trial if this method returns :obj:`True`, and to wait for human evaluation if this method returns :obj:`False`. """ - return len(self.active_trials) < self.n_generate + return len(get_active_trials(self._study._study_id, self._study._storage)) < self.n_generate def get_active_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]: From 22fdc50102100a6cce39eaf6a48a2c0177a6c505 Mon Sep 17 00:00:00 2001 From: c-bata Date: Fri, 1 Sep 2023 13:21:42 +0900 Subject: [PATCH 03/26] Bump the version up to v0.13.0b1 --- optuna_dashboard/__init__.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/__init__.py b/optuna_dashboard/__init__.py index 3d363cf4..0ab0cd89 100644 --- a/optuna_dashboard/__init__.py +++ b/optuna_dashboard/__init__.py @@ -15,4 +15,4 @@ from ._note import get_note # noqa from ._note import save_note # noqa -__version__ = "0.12.0" +__version__ = "0.13.0b1" From b929bc6b36ec96ff2bf837b4247089a31ef3a272 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Fri, 1 Sep 2023 14:36:07 +0900 Subject: [PATCH 04/26] [WIP] --- optuna_dashboard/preferential/_study.py | 4 +++- python_tests/test_api.py | 5 +++-- 2 files changed, 6 insertions(+), 3 deletions(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index 57ab8e9f..8895d604 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -276,7 +276,9 @@ class PreferentialStudy: to generate a new trial if this method returns :obj:`True`, and to wait for human evaluation if this method returns :obj:`False`. """ - return len(get_active_trials(self._study._study_id, self._study._storage)) < self.n_generate + return ( + len(get_active_trials(self._study._study_id, self._study._storage)) < self.n_generate + ) def get_active_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]: diff --git a/python_tests/test_api.py b/python_tests/test_api.py index ae50e29a..fc67e12e 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -108,6 +108,8 @@ class APITestCase(TestCase): study.mark_comparison_ready(trial) study.report_preference(study.trials[0], study.trials[1]) + assert len(study.best_trials) == 1 + app = create_app(storage) study_id = study._study._study_id status, _, body = send_request( @@ -119,9 +121,8 @@ class APITestCase(TestCase): self.assertEqual(status, 200) best_trials = json.loads(body)["best_trials"] - assert len(best_trials) == 2 + assert len(best_trials) == 1 assert best_trials[0]["number"] == 0 - assert best_trials[1]["number"] == 2 def test_report_preference(self) -> None: storage = optuna.storages.InMemoryStorage() From 9c8a0fd71e454cd94a843656032af6b148e369d5 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Mon, 4 Sep 2023 19:13:46 +0900 Subject: [PATCH 05/26] Remove mark_comparison_ready --- python_tests/preferential/test_study.py | 2 +- python_tests/test_api.py | 13 +++++++------ 2 files changed, 8 insertions(+), 7 deletions(-) diff --git a/python_tests/preferential/test_study.py b/python_tests/preferential/test_study.py index abc1bc4d..7fc1aaba 100644 --- a/python_tests/preferential/test_study.py +++ b/python_tests/preferential/test_study.py @@ -263,7 +263,7 @@ def test_get_trials_state_option(storage_supplier: Callable[[], StorageSupplier] with storage_supplier() as storage: study = create_study(n_generate=4, storage=storage) for _ in range(3): - trial = study.ask() + study.ask() better, worse = study.trials[:2] study.report_preference(better, worse) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index ba3b0dc4..9bdaf4d3 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -104,7 +104,7 @@ class APITestCase(TestCase): storage = optuna.storages.InMemoryStorage() study = create_study(n_generate=4, storage=storage) for _ in range(3): - trial = study.ask() + study.ask() study.report_preference(study.trials[0], study.trials[1]) assert len(study.best_trials) == 1 @@ -127,7 +127,7 @@ class APITestCase(TestCase): storage = optuna.storages.InMemoryStorage() study = create_study(n_generate=4, storage=storage) for _ in range(3): - trial = study.ask() + study.ask() app = create_app(storage) study_id = study._study._study_id @@ -157,21 +157,22 @@ class APITestCase(TestCase): for _ in range(3): trial = study.ask() trials.append(trial) + study.report_preference(trials[0], trials[1]) + study.report_preference(trials[2], trials[1]) app = create_app(storage) study_id = study._study._study_id status, _, _ = send_request( app, - f"/api/studies/{study_id}/{trials[1]._trial_id}/skip", + f"/api/studies/{study_id}/{trials[0]._trial_id}/skip", "POST", content_type="application/json", ) self.assertEqual(status, 204) best_trials = study.best_trials - assert len(best_trials) == 2 - assert best_trials[0].number == 0 - assert best_trials[1].number == 2 + assert len(best_trials) == 1 + assert best_trials[0].number == 2 def test_create_study(self) -> None: for name, directions, expected_status in [ From d774de65c5c347ddf9aaaadc29a603fd6915e85e Mon Sep 17 00:00:00 2001 From: Contramundum Date: Mon, 4 Sep 2023 19:19:17 +0900 Subject: [PATCH 06/26] Fix mypy error --- python_tests/test_api.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 9bdaf4d3..5085b598 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -157,8 +157,8 @@ class APITestCase(TestCase): for _ in range(3): trial = study.ask() trials.append(trial) - study.report_preference(trials[0], trials[1]) - study.report_preference(trials[2], trials[1]) + study.report_preference(study.trials[0], study.trials[1]) + study.report_preference(study.trials[2], study.trials[1]) app = create_app(storage) study_id = study._study._study_id From 32ab08954dbb94dbb54eaf1c796dc93d46248c82 Mon Sep 17 00:00:00 2001 From: contramundum53 Date: Wed, 6 Sep 2023 11:04:58 +0900 Subject: [PATCH 07/26] Update optuna_dashboard/preferential/_study.py --- optuna_dashboard/preferential/_study.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index e412d0ee..702145e3 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -245,7 +245,7 @@ def get_best_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list worse_numbers = {worse for _, worse in preferences} nondominated_numbers = {better for better, _ in preferences if better not in worse_numbers} trials = storage.get_all_trials( - study_id, deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING) + study_id, deepcopy=False, ) study_system_attrs = storage.get_study_system_attrs(study_id) From 60ea9a5a59ee8507332e7fafc749ba1b11f56bcd Mon Sep 17 00:00:00 2001 From: contramundum53 Date: Wed, 6 Sep 2023 12:36:17 +0900 Subject: [PATCH 08/26] Update optuna_dashboard/preferential/_study.py Co-authored-by: c-bata --- optuna_dashboard/preferential/_study.py | 4 +--- 1 file changed, 1 insertion(+), 3 deletions(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index 702145e3..4b6f8c2c 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -244,9 +244,7 @@ def get_best_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list preferences = get_preferences(study_id, storage) worse_numbers = {worse for _, worse in preferences} nondominated_numbers = {better for better, _ in preferences if better not in worse_numbers} - trials = storage.get_all_trials( - study_id, deepcopy=False, - ) + trials = storage.get_all_trials(study_id, deepcopy=False) study_system_attrs = storage.get_study_system_attrs(study_id) From e39357bd20a89d804507019a091b3fdc077cf057 Mon Sep 17 00:00:00 2001 From: c-bata Date: Wed, 6 Sep 2023 18:01:32 +0900 Subject: [PATCH 09/26] Support user-defined plotly figures --- optuna_dashboard/__init__.py | 1 + optuna_dashboard/_app.py | 4 + optuna_dashboard/_custom_plot_data.py | 115 ++++++++++++++++++ optuna_dashboard/_serializer.py | 5 + optuna_dashboard/ts/apiClient.ts | 2 + .../ts/components/StudyHistory.tsx | 12 ++ .../ts/components/UserDefinedPlot.tsx | 16 +++ optuna_dashboard/ts/types/index.d.ts | 6 + pyproject.toml | 1 + python_tests/test_custom_plot_data.py | 63 ++++++++++ python_tests/test_serializers.py | 4 +- 11 files changed, 227 insertions(+), 2 deletions(-) create mode 100644 optuna_dashboard/_custom_plot_data.py create mode 100644 optuna_dashboard/ts/components/UserDefinedPlot.tsx create mode 100644 python_tests/test_custom_plot_data.py diff --git a/optuna_dashboard/__init__.py b/optuna_dashboard/__init__.py index 3d363cf4..493af736 100644 --- a/optuna_dashboard/__init__.py +++ b/optuna_dashboard/__init__.py @@ -1,5 +1,6 @@ from ._app import run_server # noqa from ._app import wsgi # noqa +from ._custom_plot_data import save_plotly_graph_object # noqa from ._form_widget import ChoiceWidget # noqa from ._form_widget import dict_to_form_widget # noqa from ._form_widget import ObjectiveChoiceWidget # noqa diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 5f433072..c32c2061 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -25,6 +25,7 @@ from . import _note as note from ._bottle_util import BottleViewReturn from ._bottle_util import json_api_view from ._cached_extra_study_property import get_cached_extra_study_property +from ._custom_plot_data import get_plotly_graph_objects from ._importance import get_param_importance_from_trials_cache from ._pareto_front import get_pareto_front_trials from ._preferential_history import NewHistory @@ -214,6 +215,8 @@ def create_app( union_user_attrs, has_intermediate_values, ) = get_cached_extra_study_property(study_id, trials) + + plotly_graph_objects = get_plotly_graph_objects(system_attrs) return serialize_study_detail( summary, best_trials, @@ -222,6 +225,7 @@ def create_app( union, union_user_attrs, has_intermediate_values, + plotly_graph_objects, ) @app.get("/api/studies//param_importances") diff --git a/optuna_dashboard/_custom_plot_data.py b/optuna_dashboard/_custom_plot_data.py new file mode 100644 index 00000000..d669d4ad --- /dev/null +++ b/optuna_dashboard/_custom_plot_data.py @@ -0,0 +1,115 @@ +from __future__ import annotations + +import math +from typing import TYPE_CHECKING +import uuid + +from optuna import Study + + +if TYPE_CHECKING: + from typing import Any + + from optuna.storages import BaseStorage + import plotly.graph_objs as go + + +SYSTEM_ATTR_PLOT_DATA = "dashboard:plot_data:" +SYSTEM_ATTR_MAX_LENGTH = 2045 + + +def save_plotly_graph_object( + study: Study, figure: go.Figure, *, graph_object_id: str | None = None +) -> str: + """Save the user-defined plotly's graph object to the study. + + Example: + + .. code-block:: python + + import optuna + from optuna_dashboard import save_plotly_graph_object + + def objective(trial): + x = trial.suggest_float("x", -100, 100) + y = trial.suggest_categorical("y", [-1, 0, 1]) + return x**2 + y + + study = optuna.create_study() + study.optimize(objective, n_trials=100) + + figure = optuna.visualization.plot_optimization_history(study) + save_plotly_graph_object(study, figure) + + Args: + study: + Target study object. + plot_data: + The plotly's graph object to save. + graph_object_id: + Unique identifier of the graph object. If specified, the graph object is overwritten. + + Returns: + The graph object ID. + """ + storage = study._storage + study_id = study._study_id + + graph_object_id = graph_object_id or str(uuid.uuid4()) + key = SYSTEM_ATTR_PLOT_DATA + graph_object_id + ":" + plot_data_json_str = figure.to_json() + save_graph_object_json(storage, study_id, key, plot_data_json_str) + return graph_object_id + + +def save_graph_object_json( + storage: BaseStorage, study_id: int, key_prefix: str, plot_data_json_str: str +) -> None: + plot_data_system_attrs = split_plot_data(plot_data_json_str, key_prefix) + for k, v in plot_data_system_attrs.items(): + storage.set_study_system_attr(study_id, k, v) + + # Clear previous graph object attributes + study_system_attrs = storage.get_study_system_attrs(study_id) + all_plot_data_system_attrs = [k for k in study_system_attrs if k.startswith(key_prefix)] + if len(all_plot_data_system_attrs) > len(plot_data_system_attrs): + for i in range(len(plot_data_system_attrs), len(all_plot_data_system_attrs)): + storage.set_study_system_attr(study_id, f"{key_prefix}{i}", "") + + +def list_graph_object_ids(system_attrs: dict[str, Any]) -> list[str]: + titles = set() + for key in system_attrs: + if not key.startswith(SYSTEM_ATTR_PLOT_DATA): + continue + + s = key.split(":", maxsplit=2) # e.g. ["dashboard", "plot_data", "Optimization History:1"] + if len(s) != 3: + continue + # Please note that title may contain ":". + title = s[2].rsplit(":", maxsplit=1)[0] + titles.add(title) + return list(titles) + + +def get_plotly_graph_objects(system_attrs: dict[str, Any]) -> dict[str, str]: + graph_objects = {} + for title in list_graph_object_ids(system_attrs): + key_prefix = SYSTEM_ATTR_PLOT_DATA + title + ":" + plot_data_attrs = {k: v for k, v in system_attrs.items() if k.startswith(key_prefix)} + graph_objects[title] = concat_plot_data(plot_data_attrs, key_prefix) + return graph_objects + + +def split_plot_data(plot_data_str: str, key_prefix: str) -> dict[str, str]: + plot_data_len = len(plot_data_str) + attrs = {} + for i in range(math.ceil(plot_data_len / SYSTEM_ATTR_MAX_LENGTH)): + start = i * SYSTEM_ATTR_MAX_LENGTH + end = min((i + 1) * SYSTEM_ATTR_MAX_LENGTH, plot_data_len) + attrs[f"{key_prefix}{i}"] = plot_data_str[start:end] + return attrs + + +def concat_plot_data(plot_data_attrs: dict[str, str], key_prefix: str) -> str: + return "".join(plot_data_attrs[f"{key_prefix}{i}"] for i in range(len(plot_data_attrs))) diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 06b53c42..19acbbd5 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -132,6 +132,7 @@ def serialize_study_detail( union: list[tuple[str, BaseDistribution]], union_user_attrs: list[tuple[str, bool]], has_intermediate_values: bool, + plotly_graph_objects: dict[str, str], ) -> dict[str, Any]: serialized: dict[str, Any] = { "name": summary.study_name, @@ -162,6 +163,10 @@ def serialize_study_detail( serialized["form_widgets"] = form_widgets if serialized["is_preferential"]: serialized["preference_history"] = serialize_preference_history(system_attrs) + serialized["plotly_graph_objects"] = [ + {"id": id_, "graph_object": graph_object} + for id_, graph_object in plotly_graph_objects.items() + ] return serialized diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index e23fc2ff..e62e0e42 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -93,6 +93,7 @@ interface StudyDetailResponse { objective_names?: string[] form_widgets?: FormWidgets preference_history?: PreferenceHistoryResponce[] + plotly_graph_objects: PlotlyGraphObject[] } export const getStudyDetailAPI = ( @@ -131,6 +132,7 @@ export const getStudyDetailAPI = ( preference_history: res.data.preference_history?.map( convertPreferenceHistory ), + plotly_graph_objects: res.data.plotly_graph_objects, } }) } diff --git a/optuna_dashboard/ts/components/StudyHistory.tsx b/optuna_dashboard/ts/components/StudyHistory.tsx index b47c557a..5cca3671 100644 --- a/optuna_dashboard/ts/components/StudyHistory.tsx +++ b/optuna_dashboard/ts/components/StudyHistory.tsx @@ -15,6 +15,7 @@ import { GraphIntermediateValues } from "./GraphIntermediateValues" import Grid2 from "@mui/material/Unstable_Grid2" import { DataGrid, DataGridColumn } from "./DataGrid" import { GraphHyperparameterImportance } from "./GraphHyperparameterImportances" +import { UserDefinedPlot } from "./UserDefinedPlot" import { BestTrialsCard } from "./BestTrialsCard" import { useStudyDetailValue, @@ -102,6 +103,17 @@ export const StudyHistory: FC<{ studyId: number }> = ({ studyId }) => { /> + {studyDetail !== null && + studyDetail.plotly_graph_objects.map((go) => ( + + + + ))} {studyDetail !== null && studyDetail.directions.length == 1 && diff --git a/optuna_dashboard/ts/components/UserDefinedPlot.tsx b/optuna_dashboard/ts/components/UserDefinedPlot.tsx new file mode 100644 index 00000000..2a6f98db --- /dev/null +++ b/optuna_dashboard/ts/components/UserDefinedPlot.tsx @@ -0,0 +1,16 @@ +import * as plotly from "plotly.js-dist-min" +import React, { FC, useEffect } from "react" +import { Box } from "@mui/material" + +export const UserDefinedPlot: FC<{ + graphObject: PlotlyGraphObject +}> = ({ graphObject }) => { + const plotDomId = `user-defined-plot:${graphObject.id}` + + useEffect(() => { + const parsed = JSON.parse(graphObject.graph_object) + plotly.react(plotDomId, parsed.data, parsed.layout) + }, [graphObject]) + + return +} diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index 646d64cf..b7b35797 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -182,6 +182,11 @@ type FormWidgets = widgets: UserAttrFormWidget[] } +type PlotlyGraphObject = { + id: string + graph_object: string +} + type StudyDetail = { id: number name: string @@ -199,6 +204,7 @@ type StudyDetail = { objective_names?: string[] form_widgets?: FormWidgets preference_history?: PreferenceHistory[] + plotly_graph_objects: PlotlyGraphObject[] } type StudyDetails = { diff --git a/pyproject.toml b/pyproject.toml index 0555f701..9b7ce731 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -43,6 +43,7 @@ docs = [ test = [ "coverage", + "plotly", "pytest", "moto[s3]", ] diff --git a/python_tests/test_custom_plot_data.py b/python_tests/test_custom_plot_data.py new file mode 100644 index 00000000..fdf738d7 --- /dev/null +++ b/python_tests/test_custom_plot_data.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +from unittest.mock import patch + +import optuna +from optuna_dashboard import _custom_plot_data as custom_plot_data +from optuna_dashboard import save_plotly_graph_object + + +def get_dummy_study() -> optuna.Study: + def objective(trial: optuna.Trial) -> float: + x = trial.suggest_float("x", -100, 100) + y = trial.suggest_categorical("y", [-1, 0, 1]) + return x**2 + y + + study = optuna.create_study() + optuna.logging.set_verbosity(optuna.logging.ERROR) + study.optimize(objective, n_trials=100) + return study + + +def test_save_plotly_graph_object() -> None: + # Save history plot + dummy_study = get_dummy_study() + plot_data = optuna.visualization.plot_optimization_history(dummy_study) + graph_object_id = save_plotly_graph_object(dummy_study, plot_data) + + study_system_attrs = dummy_study._storage.get_study_system_attrs(dummy_study._study_id) + plot_data_dict = custom_plot_data.get_plotly_graph_objects(study_system_attrs) + assert len(plot_data_dict) == 1 + assert plot_data_dict[graph_object_id] == plot_data.to_json() + + # Save parallel coordinate plot + plot_data = optuna.visualization.plot_parallel_coordinate(dummy_study) + graph_object_id = save_plotly_graph_object(dummy_study, plot_data) + + study_system_attrs = dummy_study._storage.get_study_system_attrs(dummy_study._study_id) + plot_data_dict = custom_plot_data.get_plotly_graph_objects(study_system_attrs) + assert len(plot_data_dict) == 2 + assert plot_data_dict[graph_object_id] == plot_data.to_json() + + +def test_update_plotly_graph_object() -> None: + # Save history plot + dummy_study = get_dummy_study() + plot_data = optuna.visualization.plot_optimization_history(dummy_study) + graph_object_id = save_plotly_graph_object(dummy_study, plot_data) + + study_system_attrs = dummy_study._storage.get_study_system_attrs(dummy_study._study_id) + plot_data_dict = custom_plot_data.get_plotly_graph_objects(study_system_attrs) + assert len(plot_data_dict) == 1 + assert plot_data_dict[graph_object_id] == plot_data.to_json() + + # Save parallel coordinate plot + plot_data = optuna.visualization.plot_parallel_coordinate(dummy_study) + graph_object_id = save_plotly_graph_object( + dummy_study, plot_data, graph_object_id=graph_object_id + ) + + study_system_attrs = dummy_study._storage.get_study_system_attrs(dummy_study._study_id) + plot_data_dict = custom_plot_data.get_plotly_graph_objects(study_system_attrs) + assert len(plot_data_dict) == 1 + assert plot_data_dict[graph_object_id] == plot_data.to_json() diff --git a/python_tests/test_serializers.py b/python_tests/test_serializers.py index a90e0de7..72db7b26 100644 --- a/python_tests/test_serializers.py +++ b/python_tests/test_serializers.py @@ -29,7 +29,7 @@ def test_get_study_detail_is_preferential() -> None: assert len(study_summaries) == 1 study_summary = study_summaries[0] - study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False) + study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False, {}) assert study_detail["is_preferential"] @@ -40,7 +40,7 @@ def test_get_study_detail_is_not_preferential() -> None: assert len(study_summaries) == 1 study_summary = study_summaries[0] - study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False) + study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False, {}) assert not study_detail["is_preferential"] From 5b64b38bf21a9d5dd157d529d70fe17cdb900ed6 Mon Sep 17 00:00:00 2001 From: c-bata Date: Wed, 6 Sep 2023 18:15:24 +0900 Subject: [PATCH 10/26] Fix flake8 error --- python_tests/test_custom_plot_data.py | 2 -- 1 file changed, 2 deletions(-) diff --git a/python_tests/test_custom_plot_data.py b/python_tests/test_custom_plot_data.py index fdf738d7..4d01af96 100644 --- a/python_tests/test_custom_plot_data.py +++ b/python_tests/test_custom_plot_data.py @@ -1,7 +1,5 @@ from __future__ import annotations -from unittest.mock import patch - import optuna from optuna_dashboard import _custom_plot_data as custom_plot_data from optuna_dashboard import save_plotly_graph_object From 7c5e0a84e0cf3db9e5ea5ca3d7c53c18ee82d3c8 Mon Sep 17 00:00:00 2001 From: c-bata Date: Wed, 6 Sep 2023 18:38:29 +0900 Subject: [PATCH 11/26] Update docs --- docs/api.rst | 1 + 1 file changed, 1 insertion(+) diff --git a/docs/api.rst b/docs/api.rst index 09a9c14b..aadd2718 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -14,6 +14,7 @@ General APIs optuna_dashboard.wsgi optuna_dashboard.set_objective_names optuna_dashboard.save_note + optuna_dashboard.save_plotly_graph_object Human-in-the-loop ----------------- From 5f1d3bb2c416e765dd45bc77afc43893950dff0e Mon Sep 17 00:00:00 2001 From: c-bata Date: Thu, 7 Sep 2023 10:31:10 +0900 Subject: [PATCH 12/26] Validate graph object id --- optuna_dashboard/_custom_plot_data.py | 21 +++++++++++++++++++++ python_tests/test_custom_plot_data.py | 25 +++++++++++++++++++++++++ 2 files changed, 46 insertions(+) diff --git a/optuna_dashboard/_custom_plot_data.py b/optuna_dashboard/_custom_plot_data.py index d669d4ad..a88fc0a4 100644 --- a/optuna_dashboard/_custom_plot_data.py +++ b/optuna_dashboard/_custom_plot_data.py @@ -48,10 +48,14 @@ def save_plotly_graph_object( The plotly's graph object to save. graph_object_id: Unique identifier of the graph object. If specified, the graph object is overwritten. + This must be a valid HTML id attribute value. Returns: The graph object ID. """ + if graph_object_id is not None and not is_valid_html_name(graph_object_id): + raise ValueError("graph_object_id must be a valid HTML id attribute value.") + storage = study._storage study_id = study._study_id @@ -113,3 +117,20 @@ def split_plot_data(plot_data_str: str, key_prefix: str) -> dict[str, str]: def concat_plot_data(plot_data_attrs: dict[str, str], key_prefix: str) -> str: return "".join(plot_data_attrs[f"{key_prefix}{i}"] for i in range(len(plot_data_attrs))) + + +def is_valid_html_name(graph_object_id: str) -> bool: + if len(graph_object_id) == 0: + return False + + # Must begin with a letter [A-Za-z] + if not ("a" <= graph_object_id[0] <= "z" or "A" <= graph_object_id[0] <= "Z"): + return False + + # Can only contain letters [A-Za-z], numbers [0-9], hyphens ("-"), underscores ("_"), colons, and periods. + if not all( + "a" <= c <= "z" or "A" <= c <= "Z" or "0" <= c <= "9" or c in ("-", "_", ":", ".") + for c in graph_object_id[1:] + ): + return False + return True diff --git a/python_tests/test_custom_plot_data.py b/python_tests/test_custom_plot_data.py index 4d01af96..f73f7097 100644 --- a/python_tests/test_custom_plot_data.py +++ b/python_tests/test_custom_plot_data.py @@ -3,6 +3,7 @@ from __future__ import annotations import optuna from optuna_dashboard import _custom_plot_data as custom_plot_data from optuna_dashboard import save_plotly_graph_object +import pytest def get_dummy_study() -> optuna.Study: @@ -59,3 +60,27 @@ def test_update_plotly_graph_object() -> None: plot_data_dict = custom_plot_data.get_plotly_graph_objects(study_system_attrs) assert len(plot_data_dict) == 1 assert plot_data_dict[graph_object_id] == plot_data.to_json() + + +@pytest.mark.parametrize( + "name", + [ + "a", + "a1-:_.", + ], +) +def test_is_valid_html_name(name): + assert custom_plot_data.is_valid_html_name(name) + + +@pytest.mark.parametrize( + "name", + [ + "0", + "a,", + "a b", + "aあいうえお", + ], +) +def test_is_invalid_html_name(name): + assert not custom_plot_data.is_valid_html_name(name) From 0caa0e6ac4c2c57325a7fe7d7b221751670e3255 Mon Sep 17 00:00:00 2001 From: c-bata Date: Thu, 7 Sep 2023 10:40:02 +0900 Subject: [PATCH 13/26] Make UserDefinedPlot half widths --- .../ts/components/StudyHistory.tsx | 21 +++++++++---------- 1 file changed, 10 insertions(+), 11 deletions(-) diff --git a/optuna_dashboard/ts/components/StudyHistory.tsx b/optuna_dashboard/ts/components/StudyHistory.tsx index 5cca3671..907acd1a 100644 --- a/optuna_dashboard/ts/components/StudyHistory.tsx +++ b/optuna_dashboard/ts/components/StudyHistory.tsx @@ -103,17 +103,6 @@ export const StudyHistory: FC<{ studyId: number }> = ({ studyId }) => { /> - {studyDetail !== null && - studyDetail.plotly_graph_objects.map((go) => ( - - - - ))} {studyDetail !== null && studyDetail.directions.length == 1 && @@ -136,6 +125,16 @@ export const StudyHistory: FC<{ studyId: number }> = ({ studyId }) => { + {studyDetail !== null && + studyDetail.plotly_graph_objects.map((go) => ( + + + + + + + + ))} From e3b7e1d89360357b7a353dd20ad3b62338305711 Mon Sep 17 00:00:00 2001 From: c-bata Date: Thu, 7 Sep 2023 10:48:11 +0900 Subject: [PATCH 14/26] Avoid to crash the whole page when given invalid figures --- optuna_dashboard/ts/components/UserDefinedPlot.tsx | 9 +++++++-- 1 file changed, 7 insertions(+), 2 deletions(-) diff --git a/optuna_dashboard/ts/components/UserDefinedPlot.tsx b/optuna_dashboard/ts/components/UserDefinedPlot.tsx index 2a6f98db..029c4b57 100644 --- a/optuna_dashboard/ts/components/UserDefinedPlot.tsx +++ b/optuna_dashboard/ts/components/UserDefinedPlot.tsx @@ -8,8 +8,13 @@ export const UserDefinedPlot: FC<{ const plotDomId = `user-defined-plot:${graphObject.id}` useEffect(() => { - const parsed = JSON.parse(graphObject.graph_object) - plotly.react(plotDomId, parsed.data, parsed.layout) + try { + const parsed = JSON.parse(graphObject.graph_object) + plotly.react(plotDomId, parsed.data, parsed.layout) + } catch (e) { + // Avoid to crash the whole page when given invalid grpah objects. + console.error(e) + } }, [graphObject]) return From 167a4de6c4f16a7d61cdf67c371feef74d8acbce Mon Sep 17 00:00:00 2001 From: c-bata Date: Thu, 7 Sep 2023 10:53:08 +0900 Subject: [PATCH 15/26] Fix tests --- optuna_dashboard/_custom_plot_data.py | 12 +++++------- python_tests/test_custom_plot_data.py | 10 +++++----- 2 files changed, 10 insertions(+), 12 deletions(-) diff --git a/optuna_dashboard/_custom_plot_data.py b/optuna_dashboard/_custom_plot_data.py index a88fc0a4..a4dfa4af 100644 --- a/optuna_dashboard/_custom_plot_data.py +++ b/optuna_dashboard/_custom_plot_data.py @@ -53,7 +53,7 @@ def save_plotly_graph_object( Returns: The graph object ID. """ - if graph_object_id is not None and not is_valid_html_name(graph_object_id): + if graph_object_id is not None and not is_valid_graph_object_id(graph_object_id): raise ValueError("graph_object_id must be a valid HTML id attribute value.") storage = study._storage @@ -119,18 +119,16 @@ def concat_plot_data(plot_data_attrs: dict[str, str], key_prefix: str) -> str: return "".join(plot_data_attrs[f"{key_prefix}{i}"] for i in range(len(plot_data_attrs))) -def is_valid_html_name(graph_object_id: str) -> bool: +def is_valid_graph_object_id(graph_object_id: str) -> bool: if len(graph_object_id) == 0: return False - # Must begin with a letter [A-Za-z] - if not ("a" <= graph_object_id[0] <= "z" or "A" <= graph_object_id[0] <= "Z"): - return False - - # Can only contain letters [A-Za-z], numbers [0-9], hyphens ("-"), underscores ("_"), colons, and periods. + # Can only contain letters [A-Za-z], numbers [0-9], hyphens ("-"), underscores ("_"), + # colons, and periods. if not all( "a" <= c <= "z" or "A" <= c <= "Z" or "0" <= c <= "9" or c in ("-", "_", ":", ".") for c in graph_object_id[1:] ): return False + # Unlike HTML id attribute, graph object id can begin with a letter [A-Za-z] return True diff --git a/python_tests/test_custom_plot_data.py b/python_tests/test_custom_plot_data.py index f73f7097..3dcfc856 100644 --- a/python_tests/test_custom_plot_data.py +++ b/python_tests/test_custom_plot_data.py @@ -65,22 +65,22 @@ def test_update_plotly_graph_object() -> None: @pytest.mark.parametrize( "name", [ + "0", "a", "a1-:_.", ], ) -def test_is_valid_html_name(name): - assert custom_plot_data.is_valid_html_name(name) +def test_is_valid_graph_object_id(name: str) -> None: + assert custom_plot_data.is_valid_graph_object_id(name) @pytest.mark.parametrize( "name", [ - "0", "a,", "a b", "aあいうえお", ], ) -def test_is_invalid_html_name(name): - assert not custom_plot_data.is_valid_html_name(name) +def test_is_invalid_graph_object_id(name: str) -> None: + assert not custom_plot_data.is_valid_graph_object_id(name) From 7b0dce75aabe0ad254b14a1bf7bdf9a01fc72932 Mon Sep 17 00:00:00 2001 From: keisuke umezawa Date: Thu, 7 Sep 2023 13:52:57 +0900 Subject: [PATCH 16/26] Update python-coverage.yml --- .github/workflows/python-coverage.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/python-coverage.yml b/.github/workflows/python-coverage.yml index 4bda6350..c96f506b 100644 --- a/.github/workflows/python-coverage.yml +++ b/.github/workflows/python-coverage.yml @@ -45,4 +45,4 @@ jobs: with: token: ${{ secrets.CODECOV_TOKEN }} file: ./coverage.xml - fail_ci_if_error: true + fail_ci_if_error: false From 5622eb58a963ab173d876b221af593424b1e5d8c Mon Sep 17 00:00:00 2001 From: Contramundum Date: Thu, 7 Sep 2023 15:31:59 +0900 Subject: [PATCH 17/26] Fix UI to adapt to best_trial change --- optuna_dashboard/preferential/_study.py | 22 ++++++++++++++++--- .../preferential/_system_attrs.py | 16 +++++++++----- optuna_dashboard/preferential/samplers/gp.py | 2 +- .../ts/components/PreferentialTrials.tsx | 12 ++++++---- .../preferential/test_system_attrs.py | 7 +++--- 5 files changed, 42 insertions(+), 17 deletions(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index 7da363b2..fb54ed36 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -14,6 +14,7 @@ from optuna.trial import FrozenTrial from optuna.trial import TrialState from optuna_dashboard.preferential._system_attrs import get_n_generate from optuna_dashboard.preferential._system_attrs import get_preferences +from optuna_dashboard.preferential._system_attrs import get_skipped_trial_ids from optuna_dashboard.preferential._system_attrs import is_skipped_trial from optuna_dashboard.preferential._system_attrs import report_preferences from optuna_dashboard.preferential._system_attrs import set_n_generate @@ -243,8 +244,11 @@ class PreferentialStudy: Returns: A list of the pair of FrozenTrial objects. The left trial is better than the right one. """ + + preferences = get_preferences( + self._study._storage.get_study_system_attrs(self._study._study_id) + ) # Must come before study.get_trials() trials = self._study.get_trials(deepcopy=deepcopy) - preferences = get_preferences(self._study._study_id, self._study._storage) return [(trials[better], trials[worse]) for (better, worse) in preferences] def set_user_attr(self, key: str, value: Any) -> None: @@ -269,11 +273,23 @@ class PreferentialStudy: to generate a new trial if this method returns :obj:`True`, and to wait for human evaluation if this method returns :obj:`False`. """ - return len(self.best_trials) < get_n_generate(self._study.system_attrs) + study_system_attrs = self._study._storage.get_study_system_attrs( + self._study._study_id + ) # Must come before _study.get_trials() + trials = self._study.get_trials(deepcopy=False) + + all_trial_ids = {t._trial_id for t in trials} + bad_trial_ids = { + trials[worse]._trial_id for (_, worse) in get_preferences(study_system_attrs) + } + skipped_trial_ids = set(get_skipped_trial_ids(study_system_attrs)) + + active_trial_ids = all_trial_ids - bad_trial_ids - skipped_trial_ids + return len(active_trial_ids) < get_n_generate(self._study.system_attrs) def get_best_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]: - preferences = get_preferences(study_id, storage) + preferences = get_preferences(storage.get_study_system_attrs(study_id)) worse_numbers = {worse for _, worse in preferences} nondominated_numbers = {better for better, _ in preferences if better not in worse_numbers} trials = storage.get_all_trials(study_id, deepcopy=False) diff --git a/optuna_dashboard/preferential/_system_attrs.py b/optuna_dashboard/preferential/_system_attrs.py index 4cb0e288..2f951594 100644 --- a/optuna_dashboard/preferential/_system_attrs.py +++ b/optuna_dashboard/preferential/_system_attrs.py @@ -35,13 +35,9 @@ def report_preferences( return preference_id -def get_preferences( - study_id: int, - storage: BaseStorage, -) -> list[tuple[int, int]]: +def get_preferences(study_system_attrs: dict[str, Any]) -> list[tuple[int, int]]: preferences: list[tuple[int, int]] = [] - system_attrs = storage.get_study_system_attrs(study_id) - for k, v in system_attrs.items(): + for k, v in study_system_attrs.items(): if not k.startswith(_SYSTEM_ATTR_PREFIX_PREFERENCE): continue preferences.extend(v) # type: ignore @@ -65,6 +61,14 @@ def is_skipped_trial(trial_id: int, study_system_attrs: dict[str, Any]) -> bool: return key in study_system_attrs +def get_skipped_trial_ids(study_system_attrs: dict[str, Any]) -> list[int]: + return [ + int(k[len(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL):]) + for k in study_system_attrs.keys() + if k.startswith(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL) + ] + + def get_n_generate(study_system_attrs: dict[str, Any]) -> int: return study_system_attrs[_SYSTEM_ATTR_N_GENERATE] diff --git a/optuna_dashboard/preferential/samplers/gp.py b/optuna_dashboard/preferential/samplers/gp.py index ff4001c1..b90a1de4 100644 --- a/optuna_dashboard/preferential/samplers/gp.py +++ b/optuna_dashboard/preferential/samplers/gp.py @@ -342,7 +342,7 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler): if len(search_space) == 0: return {} - preferences = get_preferences(study._study_id, study._storage) + preferences = get_preferences(study.system_attrs) trials = study.get_trials(deepcopy=False) if len(preferences) == 0: return {} diff --git a/optuna_dashboard/ts/components/PreferentialTrials.tsx b/optuna_dashboard/ts/components/PreferentialTrials.tsx index 93ed2f5a..27873c45 100644 --- a/optuna_dashboard/ts/components/PreferentialTrials.tsx +++ b/optuna_dashboard/ts/components/PreferentialTrials.tsx @@ -182,11 +182,15 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ return null } const theme = useTheme() + + const running_trials = studyDetail.trials.filter((t) => t.state === "Running") + const active_trials = running_trials.concat(studyDetail.best_trials) + const [displayTrials, setDisplayTrials] = useState({ - numbers: studyDetail.best_trials.map((t) => t.number), - last_number: Math.max(...studyDetail.best_trials.map((t) => t.number), -1), + numbers: active_trials.map((t) => t.number), + last_number: Math.max(...active_trials.map((t) => t.number), -1), }) - const new_trails = studyDetail.best_trials.filter( + const new_trails = active_trials.filter( (t) => displayTrials.last_number < t.number && displayTrials.numbers.find((n) => n === t.number) === undefined @@ -239,7 +243,7 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ {displayTrials.numbers.map((t, index) => ( trial.number === t)} + trial={active_trials.find((trial) => trial.number === t)} candidates={displayTrials.numbers.filter((n) => n !== -1)} hideTrial={() => { hideTrial(t) diff --git a/python_tests/preferential/test_system_attrs.py b/python_tests/preferential/test_system_attrs.py index 10448d48..34f93200 100644 --- a/python_tests/preferential/test_system_attrs.py +++ b/python_tests/preferential/test_system_attrs.py @@ -18,12 +18,13 @@ def test_report_and_get_preferences(storage_supplier: Callable[[], StorageSuppli study.ask() study_id = study._study_id - assert len(get_preferences(study_id, storage)) == 0 + + assert len(get_preferences(storage.get_study_system_attrs(study_id))) == 0 better, worse = study.trials[0], study.trials[1] report_preferences(study_id, storage, [(better.number, worse.number)]) - assert len(get_preferences(study_id, storage)) == 1 + assert len(get_preferences(storage.get_study_system_attrs(study_id))) == 1 - actual_better, actual_worse = get_preferences(study_id, storage)[0] + actual_better, actual_worse = get_preferences(storage.get_study_system_attrs(study_id))[0] assert actual_better == better.number assert actual_worse == worse.number From 94be880c3f14c824f767efe0f2e66591c9649a01 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Thu, 7 Sep 2023 15:34:22 +0900 Subject: [PATCH 18/26] Run formatter --- optuna_dashboard/ts/components/PreferentialTrials.tsx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/ts/components/PreferentialTrials.tsx b/optuna_dashboard/ts/components/PreferentialTrials.tsx index 27873c45..3bdb0cce 100644 --- a/optuna_dashboard/ts/components/PreferentialTrials.tsx +++ b/optuna_dashboard/ts/components/PreferentialTrials.tsx @@ -182,7 +182,7 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ return null } const theme = useTheme() - + const running_trials = studyDetail.trials.filter((t) => t.state === "Running") const active_trials = running_trials.concat(studyDetail.best_trials) From 3f00b42fa7dd006b504f1d2532aadc0e6e716bce Mon Sep 17 00:00:00 2001 From: contramundum53 Date: Thu, 7 Sep 2023 15:55:56 +0900 Subject: [PATCH 19/26] Update optuna_dashboard/preferential/_system_attrs.py Co-authored-by: c-bata --- optuna_dashboard/preferential/_system_attrs.py | 15 ++++++++++----- 1 file changed, 10 insertions(+), 5 deletions(-) diff --git a/optuna_dashboard/preferential/_system_attrs.py b/optuna_dashboard/preferential/_system_attrs.py index 2f951594..0483eeb7 100644 --- a/optuna_dashboard/preferential/_system_attrs.py +++ b/optuna_dashboard/preferential/_system_attrs.py @@ -62,11 +62,16 @@ def is_skipped_trial(trial_id: int, study_system_attrs: dict[str, Any]) -> bool: def get_skipped_trial_ids(study_system_attrs: dict[str, Any]) -> list[int]: - return [ - int(k[len(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL):]) - for k in study_system_attrs.keys() - if k.startswith(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL) - ] + skipped_trial_ids: list[int] = [] + for k in study_system_attrs: + if not k.startswith(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL): + continue + try: + trial_id = int(k[len(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL):]) + skipped_trial_ids.append(trial_id) + except ValueError as e: + continue + return skipped_trial_ids def get_n_generate(study_system_attrs: dict[str, Any]) -> int: From 49c49623d97efa52b86964867078f953056631a0 Mon Sep 17 00:00:00 2001 From: contramundum53 Date: Thu, 7 Sep 2023 15:56:39 +0900 Subject: [PATCH 20/26] Update optuna_dashboard/ts/components/PreferentialTrials.tsx Co-authored-by: c-bata --- optuna_dashboard/ts/components/PreferentialTrials.tsx | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/optuna_dashboard/ts/components/PreferentialTrials.tsx b/optuna_dashboard/ts/components/PreferentialTrials.tsx index 3bdb0cce..c9ab1385 100644 --- a/optuna_dashboard/ts/components/PreferentialTrials.tsx +++ b/optuna_dashboard/ts/components/PreferentialTrials.tsx @@ -183,8 +183,8 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ } const theme = useTheme() - const running_trials = studyDetail.trials.filter((t) => t.state === "Running") - const active_trials = running_trials.concat(studyDetail.best_trials) + const runningTrials = studyDetail.trials.filter((t) => t.state === "Running") + const activeTrials = running_trials.concat(studyDetail.best_trials) const [displayTrials, setDisplayTrials] = useState({ numbers: active_trials.map((t) => t.number), From 8d5c306ddd16cf33887a202672f6cde0f6641c81 Mon Sep 17 00:00:00 2001 From: contramundum53 Date: Thu, 7 Sep 2023 15:57:54 +0900 Subject: [PATCH 21/26] Update optuna_dashboard/preferential/_study.py Co-authored-by: c-bata --- optuna_dashboard/preferential/_study.py | 13 ++++--------- 1 file changed, 4 insertions(+), 9 deletions(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index fb54ed36..f47f8f9d 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -276,16 +276,11 @@ class PreferentialStudy: study_system_attrs = self._study._storage.get_study_system_attrs( self._study._study_id ) # Must come before _study.get_trials() - trials = self._study.get_trials(deepcopy=False) - - all_trial_ids = {t._trial_id for t in trials} - bad_trial_ids = { - trials[worse]._trial_id for (_, worse) in get_preferences(study_system_attrs) - } + trials = self._study.get_trials(deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING)) + worse_trial_numbers = {worse for _, worse in get_preferences(study_system_attrs)} skipped_trial_ids = set(get_skipped_trial_ids(study_system_attrs)) - - active_trial_ids = all_trial_ids - bad_trial_ids - skipped_trial_ids - return len(active_trial_ids) < get_n_generate(self._study.system_attrs) + active_trials = [t for t in trials if t.number not in worse_trial_number and t._trial_id not in skipped_trial_ids] + return len(active_trials) < get_n_generate(self._study.system_attrs) def get_best_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]: From 39877452a01516c7cbd0d712b08522b3855a8f69 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Thu, 7 Sep 2023 16:02:58 +0900 Subject: [PATCH 22/26] Apply review comments --- optuna_dashboard/preferential/_study.py | 10 ++++++++-- optuna_dashboard/preferential/_system_attrs.py | 2 +- optuna_dashboard/ts/components/PreferentialTrials.tsx | 10 +++++----- 3 files changed, 14 insertions(+), 8 deletions(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index f47f8f9d..6e093683 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -276,10 +276,16 @@ class PreferentialStudy: study_system_attrs = self._study._storage.get_study_system_attrs( self._study._study_id ) # Must come before _study.get_trials() - trials = self._study.get_trials(deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING)) + trials = self._study.get_trials( + deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING) + ) worse_trial_numbers = {worse for _, worse in get_preferences(study_system_attrs)} skipped_trial_ids = set(get_skipped_trial_ids(study_system_attrs)) - active_trials = [t for t in trials if t.number not in worse_trial_number and t._trial_id not in skipped_trial_ids] + active_trials = [ + t + for t in trials + if t.number not in worse_trial_numbers and t._trial_id not in skipped_trial_ids + ] return len(active_trials) < get_n_generate(self._study.system_attrs) diff --git a/optuna_dashboard/preferential/_system_attrs.py b/optuna_dashboard/preferential/_system_attrs.py index 0483eeb7..c4454e7e 100644 --- a/optuna_dashboard/preferential/_system_attrs.py +++ b/optuna_dashboard/preferential/_system_attrs.py @@ -67,7 +67,7 @@ def get_skipped_trial_ids(study_system_attrs: dict[str, Any]) -> list[int]: if not k.startswith(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL): continue try: - trial_id = int(k[len(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL):]) + trial_id = int(k[len(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL) :]) skipped_trial_ids.append(trial_id) except ValueError as e: continue diff --git a/optuna_dashboard/ts/components/PreferentialTrials.tsx b/optuna_dashboard/ts/components/PreferentialTrials.tsx index c9ab1385..f3166627 100644 --- a/optuna_dashboard/ts/components/PreferentialTrials.tsx +++ b/optuna_dashboard/ts/components/PreferentialTrials.tsx @@ -184,13 +184,13 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ const theme = useTheme() const runningTrials = studyDetail.trials.filter((t) => t.state === "Running") - const activeTrials = running_trials.concat(studyDetail.best_trials) + const activeTrials = runningTrials.concat(studyDetail.best_trials) const [displayTrials, setDisplayTrials] = useState({ - numbers: active_trials.map((t) => t.number), - last_number: Math.max(...active_trials.map((t) => t.number), -1), + numbers: activeTrials.map((t) => t.number), + last_number: Math.max(...activeTrials.map((t) => t.number), -1), }) - const new_trails = active_trials.filter( + const new_trails = activeTrials.filter( (t) => displayTrials.last_number < t.number && displayTrials.numbers.find((n) => n === t.number) === undefined @@ -243,7 +243,7 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({ {displayTrials.numbers.map((t, index) => ( trial.number === t)} + trial={activeTrials.find((trial) => trial.number === t)} candidates={displayTrials.numbers.filter((n) => n !== -1)} hideTrial={() => { hideTrial(t) From 226bd53953475422cbd9d5f09019f619768f09f3 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Thu, 7 Sep 2023 17:37:56 +0900 Subject: [PATCH 23/26] Fix linter --- optuna_dashboard/preferential/_system_attrs.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/preferential/_system_attrs.py b/optuna_dashboard/preferential/_system_attrs.py index c4454e7e..ab469a74 100644 --- a/optuna_dashboard/preferential/_system_attrs.py +++ b/optuna_dashboard/preferential/_system_attrs.py @@ -69,7 +69,7 @@ def get_skipped_trial_ids(study_system_attrs: dict[str, Any]) -> list[int]: try: trial_id = int(k[len(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL) :]) skipped_trial_ids.append(trial_id) - except ValueError as e: + except ValueError: continue return skipped_trial_ids From 1c01ef0127190462f7173bd1aa25e4c5fadcf66d Mon Sep 17 00:00:00 2001 From: Contramundum Date: Thu, 7 Sep 2023 17:38:55 +0900 Subject: [PATCH 24/26] Fix test --- python_tests/test_api.py | 1 - python_tests/test_preferential_history.py | 1 - 2 files changed, 2 deletions(-) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 75243f30..e0c2e1c8 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -161,7 +161,6 @@ class APITestCase(TestCase): study = create_study(storage=storage, n_generate=3) for _ in range(3): trial = study.ask() - study.mark_comparison_ready(trial) app = create_app(storage) study_id = study._study._study_id diff --git a/python_tests/test_preferential_history.py b/python_tests/test_preferential_history.py index ab524b90..51c9b0f8 100644 --- a/python_tests/test_preferential_history.py +++ b/python_tests/test_preferential_history.py @@ -19,7 +19,6 @@ def test_report_and_get_choices(storage_supplier: Callable[[], StorageSupplier]) for _ in range(5): trial = study.ask() trial.suggest_float("x", 0, 1) - study.mark_comparison_ready(trial) study_id = study._study._study_id From 4e8c6051e467f974cde0a4d2f0718ac2c4c7ccea Mon Sep 17 00:00:00 2001 From: Contramundum Date: Thu, 7 Sep 2023 17:42:08 +0900 Subject: [PATCH 25/26] Apply review comment --- optuna_dashboard/ts/components/PreferentialTrials.tsx | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/optuna_dashboard/ts/components/PreferentialTrials.tsx b/optuna_dashboard/ts/components/PreferentialTrials.tsx index f3166627..08efd9a7 100644 --- a/optuna_dashboard/ts/components/PreferentialTrials.tsx +++ b/optuna_dashboard/ts/components/PreferentialTrials.tsx @@ -42,6 +42,8 @@ const PreferentialTrial: FC<{ ) } + const isBestTrial = trial.state === "Complete" + return ( true} + isBestTrial={() => isBestTrial} directions={[]} objectiveNames={[]} /> From f152ee19953d040828d42de05f7668022e4ef8c5 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Thu, 7 Sep 2023 17:45:58 +0900 Subject: [PATCH 26/26] Fix linter --- optuna_dashboard/preferential/_system_attrs.py | 2 +- python_tests/test_api.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/optuna_dashboard/preferential/_system_attrs.py b/optuna_dashboard/preferential/_system_attrs.py index ab469a74..47c2a486 100644 --- a/optuna_dashboard/preferential/_system_attrs.py +++ b/optuna_dashboard/preferential/_system_attrs.py @@ -67,7 +67,7 @@ def get_skipped_trial_ids(study_system_attrs: dict[str, Any]) -> list[int]: if not k.startswith(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL): continue try: - trial_id = int(k[len(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL) :]) + trial_id = int(k[len(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL) :]) # noqa: E203 skipped_trial_ids.append(trial_id) except ValueError: continue diff --git a/python_tests/test_api.py b/python_tests/test_api.py index e0c2e1c8..c551e3f2 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -160,7 +160,7 @@ class APITestCase(TestCase): storage = optuna.storages.InMemoryStorage() study = create_study(storage=storage, n_generate=3) for _ in range(3): - trial = study.ask() + study.ask() app = create_app(storage) study_id = study._study._study_id