From 0f52401aad615ad280ac3847c5efd92c6dda7893 Mon Sep 17 00:00:00 2001 From: c-bata Date: Wed, 16 Mar 2022 16:25:24 +0900 Subject: [PATCH 1/2] Hide IntermediateValue chart by default --- optuna_dashboard/_app.py | 11 ++++- optuna_dashboard/_intermediate_values.py | 49 +++++++++++++++++++ optuna_dashboard/_serializer.py | 2 + optuna_dashboard/static/apiClient.ts | 2 + .../static/components/StudyDetail.tsx | 5 +- optuna_dashboard/static/types/index.d.ts | 1 + 6 files changed, 67 insertions(+), 3 deletions(-) create mode 100644 optuna_dashboard/_intermediate_values.py diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 3894ff63..831e6093 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -34,6 +34,7 @@ from optuna.study import StudySummary from optuna.trial import FrozenTrial from optuna.trial import TrialState +from ._intermediate_values import has_intermediate_values from ._search_space import get_search_space from ._serializer import serialize_study_detail from ._serializer import serialize_study_summary @@ -217,9 +218,15 @@ def create_app(storage: BaseStorage) -> Bottle: if summary is None: response.status = 404 # Not found return {"reason": f"study_id={study_id} is not found"} - trials = get_trials(storage, study_id)[after:] + trials = get_trials(storage, study_id) intersection, union = get_search_space(study_id, trials) - return serialize_study_detail(summary, trials, intersection, union) + return serialize_study_detail( + summary, + trials[after:], + intersection, + union, + has_intermediate_values(study_id, trials), + ) @app.get("/api/studies//param_importances") @handle_json_api_exception diff --git a/optuna_dashboard/_intermediate_values.py b/optuna_dashboard/_intermediate_values.py new file mode 100644 index 00000000..d9b8e463 --- /dev/null +++ b/optuna_dashboard/_intermediate_values.py @@ -0,0 +1,49 @@ +import threading +from typing import Dict +from typing import List + +from optuna.trial import FrozenTrial +from optuna.trial import TrialState + + +# In-memory cache +intermediate_values_cache_lock = threading.Lock() +intermediate_values_cache: Dict[int, "_IntermediateValues"] = {} +states_of_interest = [TrialState.COMPLETE, TrialState.PRUNED] + + +def has_intermediate_values(study_id: int, trials: List[FrozenTrial]) -> bool: + with intermediate_values_cache_lock: + intermediate_values = intermediate_values_cache.get(study_id, None) + if intermediate_values is None: + intermediate_values = _IntermediateValues() + intermediate_values.update(trials) + intermediate_values_cache[study_id] = intermediate_values + return intermediate_values.has_intermediate_values + + +class _IntermediateValues: + def __init__(self) -> None: + self._cursor: int = -1 + self.has_intermediate_values: bool = False + + def update(self, trials: List[FrozenTrial]) -> None: + if self.has_intermediate_values: + return + + next_cursor = self._cursor + for trial in reversed(trials): + if self._cursor > trial.number: + break + + if not trial.state.is_finished(): + next_cursor = trial.number + + if trial.state not in states_of_interest: + continue + + current = len(trial.intermediate_values) > 0 + if current: + self.has_intermediate_values = True + return + self._cursor = next_cursor diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 8b25e655..fc155214 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -90,6 +90,7 @@ def serialize_study_detail( trials: List[FrozenTrial], intersection: List[Tuple[str, BaseDistribution]], union: List[Tuple[str, BaseDistribution]], + has_intermediate_values: bool, ) -> Dict[str, Any]: serialized: Dict[str, Any] = { "name": summary.study_name, @@ -109,6 +110,7 @@ def serialize_study_detail( serialized["intersection_search_space"] = serialize_search_space(intersection) serialized["union_search_space"] = serialize_search_space(union) + serialized["has_intermediate_values"] = has_intermediate_values return serialized diff --git a/optuna_dashboard/static/apiClient.ts b/optuna_dashboard/static/apiClient.ts index e5dba916..6dd934a4 100644 --- a/optuna_dashboard/static/apiClient.ts +++ b/optuna_dashboard/static/apiClient.ts @@ -44,6 +44,7 @@ interface StudyDetailResponse { trials: TrialResponse[] intersection_search_space: SearchSpace[] union_search_space: SearchSpace[] + has_intermediate_values: boolean } export const getStudyDetailAPI = ( @@ -70,6 +71,7 @@ export const getStudyDetailAPI = ( trials: trials, union_search_space: res.data.union_search_space, intersection_search_space: res.data.intersection_search_space, + has_intermediate_values: res.data.has_intermediate_values, } }) } diff --git a/optuna_dashboard/static/components/StudyDetail.tsx b/optuna_dashboard/static/components/StudyDetail.tsx index 25eef599..b7f4abc4 100644 --- a/optuna_dashboard/static/components/StudyDetail.tsx +++ b/optuna_dashboard/static/components/StudyDetail.tsx @@ -190,7 +190,9 @@ export const StudyDetail: FC<{ /> diff --git a/optuna_dashboard/static/types/index.d.ts b/optuna_dashboard/static/types/index.d.ts index 80268d93..4e62e683 100644 --- a/optuna_dashboard/static/types/index.d.ts +++ b/optuna_dashboard/static/types/index.d.ts @@ -77,6 +77,7 @@ declare interface StudyDetail { trials: Trial[] intersection_search_space: SearchSpace[] union_search_space: SearchSpace[] + has_intermediate_values: boolean } declare interface StudyDetails { From 31f0c3eab24462d141a16ef0a4d58c90207f2114 Mon Sep 17 00:00:00 2001 From: c-bata Date: Wed, 16 Mar 2022 16:36:06 +0900 Subject: [PATCH 2/2] Fix TrialTable.test.tsx --- frontend_tests/TrialTable.test.tsx | 1 + 1 file changed, 1 insertion(+) diff --git a/frontend_tests/TrialTable.test.tsx b/frontend_tests/TrialTable.test.tsx index 4ddd6fc3..7a0fda2d 100644 --- a/frontend_tests/TrialTable.test.tsx +++ b/frontend_tests/TrialTable.test.tsx @@ -73,6 +73,7 @@ const study_detail = { attributes: { low: -3, high: 3 }, }, ], + has_intermediate_values: false, } it("Sort TrialTable by trial number", () => {