From 1f699f2bd7b798c687ae58909f193dd2a48b7519 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sat, 19 Mar 2022 21:47:08 +0900 Subject: [PATCH 1/3] Cache hyperparameter importances --- optuna_dashboard/_app.py | 58 ++++---------------- optuna_dashboard/_importance.py | 96 +++++++++++++++++++++++++++++++++ 2 files changed, 107 insertions(+), 47 deletions(-) create mode 100644 optuna_dashboard/_importance.py diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 33d91594..ba5d2bd5 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -23,18 +23,16 @@ from bottle import request from bottle import response from bottle import run from bottle import static_file -import optuna from optuna.exceptions import DuplicatedStudyError from optuna.storages import BaseStorage from optuna.storages import RDBStorage from optuna.storages import RedisStorage -from optuna.study import Study from optuna.study import StudyDirection from optuna.study import StudySummary from optuna.trial import FrozenTrial -from optuna.trial import TrialState from . import _note as note +from ._importance import get_param_importance_from_trials_cache from ._intermediate_values import has_intermediate_values from ._search_space import get_search_space from ._serializer import serialize_study_detail @@ -179,13 +177,6 @@ def get_trials( return trials -def get_distribution_name(param_name: str, study: Study) -> str: - for trial in study.trials: - if param_name in trial.distributions: - return trial.distributions[param_name].__class__.__name__ - assert False, "Must not reach here." - - def create_app(storage: BaseStorage) -> Bottle: app = Bottle() @@ -292,51 +283,24 @@ def create_app(storage: BaseStorage) -> Bottle: # TODO(chenghuzi): add support for selecting params via query parameters. objective_id = int(request.params.get("objective_id", 0)) try: - study_name = storage.get_study_name_from_id(study_id) - study = Study(study_name=study_name, storage=storage) + n_directions = len(storage.get_study_directions(study_id)) except KeyError: - response.status = 404 # Not found + response.status = 404 # Study is not found return {"reason": f"study_id={study_id} is not found"} - - n_directions = len(study.directions) if objective_id >= n_directions: response.status = 400 # Bad request return { "reason": f"study_id={study_id} has only {n_directions} direction(s)." } - completed_trials = [ - trial for trial in study.trials if trial.state == TrialState.COMPLETE - ] - evaluator = None - params = None - - if len(completed_trials) > 0: - try: - importances = optuna.importance.get_param_importances( - study, - evaluator=evaluator, - params=params, - target=lambda t: t.values[objective_id], - ) - except ValueError as e: - response.status = 400 # Bad request - return {"reason": str(e)} - else: - importances = {} - target_name = "Objective Value" - - return { - "target_name": target_name, - "param_importances": [ - { - "name": name, - "importance": importance, - "distribution": get_distribution_name(name, study), - } - for name, importance in importances.items() - ], - } + trials = get_trials(storage, study_id) + try: + return get_param_importance_from_trials_cache( + storage, study_id, objective_id, trials + ) + except ValueError as e: + response.status = 400 # Bad request + return {"reason": str(e)} @app.put("/api/studies//note") @json_api_view diff --git a/optuna_dashboard/_importance.py b/optuna_dashboard/_importance.py new file mode 100644 index 00000000..bfcfb358 --- /dev/null +++ b/optuna_dashboard/_importance.py @@ -0,0 +1,96 @@ +import threading +from typing import Dict +from typing import List +from typing import Tuple + + +try: + from typing import TypedDict +except ImportError: + from typing_extensions import TypedDict + +import optuna +from optuna.storages import BaseStorage +from optuna.study import Study +from optuna.trial import FrozenTrial +from optuna.trial import TrialState + + +ImportanceItemType = TypedDict( + "ImportanceItemType", + { + "name": str, + "importance": float, + "distribution": str, + }, +) +ImportanceType = TypedDict( + "ImportanceType", + { + "target_name": str, + "param_importances": List[ImportanceItemType], + }, +) + +target_name = "Objective Value" +param_importance_cache_lock = threading.Lock() +# { "{study_id}:{objective_id}" : (n_completed_trials, importance) } +param_importance_cache: Dict[str, Tuple[int, ImportanceType]] = {} + + +class StudyWrapper(Study): + def __init__( + self, storage: BaseStorage, study_id: int, cached_trials: List[FrozenTrial] + ) -> None: + study_name = storage.get_study_name_from_id(study_id) + super().__init__(study_name=study_name, storage=storage) + self._cached_trials = cached_trials + + @property + def trials(self) -> List[FrozenTrial]: + return self._cached_trials + + +def get_param_importance_from_trials_cache( + storage: BaseStorage, study_id: int, objective_id: int, trials: List[FrozenTrial] +) -> ImportanceType: + n_completed_trials = len([t for t in trials if t.state == TrialState.COMPLETE]) + if n_completed_trials == 0: + return {"target_name": target_name, "param_importances": []} + + cache_key = f"{study_id}:{objective_id}" + with param_importance_cache_lock: + cache_n_trial, cache_importance = param_importance_cache.get(cache_key, [0, {}]) + if n_completed_trials == cache_n_trial: + return cache_importance + + study = StudyWrapper(storage, study_id, trials) + importance = optuna.importance.get_param_importances( + study, target=lambda t: t.values[objective_id] + ) + converted = convert_to_importance_type(importance, trials) + param_importance_cache[cache_key] = (n_completed_trials, converted) + return converted + + +def convert_to_importance_type( + importance: Dict[str, float], trials: List[FrozenTrial] +) -> ImportanceType: + return { + "target_name": target_name, + "param_importances": [ + { + "name": name, + "importance": importance, + "distribution": get_distribution_name(name, trials), + } + for name, importance in importance.items() + ], + } + + +def get_distribution_name(param_name: str, trials: List[FrozenTrial]) -> str: + for trial in trials: + if param_name in trial.distributions: + return trial.distributions[param_name].__class__.__name__ + assert False, "Must not reach here." From b415eaf6d6af8a69fa0d292b53468f5f96707fe1 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 20 Mar 2022 01:37:01 +0900 Subject: [PATCH 2/3] Fix importance cache --- optuna_dashboard/_importance.py | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/optuna_dashboard/_importance.py b/optuna_dashboard/_importance.py index bfcfb358..dd67050d 100644 --- a/optuna_dashboard/_importance.py +++ b/optuna_dashboard/_importance.py @@ -60,7 +60,9 @@ def get_param_importance_from_trials_cache( cache_key = f"{study_id}:{objective_id}" with param_importance_cache_lock: - cache_n_trial, cache_importance = param_importance_cache.get(cache_key, [0, {}]) + cache_n_trial, cache_importance = param_importance_cache.get( + cache_key, [0, {"target_name": target_name, "param_importances": []}] + ) if n_completed_trials == cache_n_trial: return cache_importance From 6ddc855e8a0c8d8d102b4e3404320b54f82afc2d Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 20 Mar 2022 01:44:02 +0900 Subject: [PATCH 3/3] Fix lint error --- optuna_dashboard/_importance.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/_importance.py b/optuna_dashboard/_importance.py index dd67050d..aef90a63 100644 --- a/optuna_dashboard/_importance.py +++ b/optuna_dashboard/_importance.py @@ -61,7 +61,7 @@ def get_param_importance_from_trials_cache( cache_key = f"{study_id}:{objective_id}" with param_importance_cache_lock: cache_n_trial, cache_importance = param_importance_cache.get( - cache_key, [0, {"target_name": target_name, "param_importances": []}] + cache_key, (0, {"target_name": target_name, "param_importances": []}) ) if n_completed_trials == cache_n_trial: return cache_importance