From e21159d40e6539bccb6b5557e4c1c6cab55e330a Mon Sep 17 00:00:00 2001 From: i23_moririn2528 Date: Tue, 15 Aug 2023 18:56:26 +0900 Subject: [PATCH] fix preferential function --- optuna_dashboard/preferential/_study.py | 33 +++++++++----- .../preferential/_system_attrs.py | 44 +++++++++---------- 2 files changed, 43 insertions(+), 34 deletions(-) diff --git a/optuna_dashboard/preferential/_study.py b/optuna_dashboard/preferential/_study.py index 9bbd6d2f..0f2c7d84 100644 --- a/optuna_dashboard/preferential/_study.py +++ b/optuna_dashboard/preferential/_study.py @@ -31,16 +31,7 @@ class PreferentialStudy: @property def best_trials(self) -> list[FrozenTrial]: - ready_trials = [ - t - for t in self._study.get_trials( - deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING) - ) - if t.system_attrs.get(_SYSTEM_ATTR_COMPARISON_READY) is True - ] - preferences = get_preferences(self._study, deepcopy=False) - worse_numbers = {worse.number for _, worse in preferences} - return [copy.deepcopy(t) for t in ready_trials if t.number not in worse_numbers] + return get_best_trials(self._study._study_id, self._study._storage) @property def study_name(self) -> str: @@ -80,10 +71,12 @@ class PreferentialStudy: if not isinstance(worse_trials, list): worse_trials = [worse_trials] - report_preferences(self._study, [(b, w) for b in better_trials for w in worse_trials]) + report_preferences(self._study._study_id, self._study._storage, [(b.number, w.number) for b in better_trials for w in worse_trials]) def get_preferences(self, *, deepcopy: bool = True) -> list[tuple[FrozenTrial, FrozenTrial]]: - return get_preferences(self._study, deepcopy=deepcopy) + trials = self._study.get_trials(deepcopy=deepcopy) + preferences = get_preferences(self._study, trials) + return [(trials[better], trials[worse]) for (better, worse) in preferences] def set_user_attr(self, key: str, value: Any) -> None: self._study.set_user_attr(key, value) @@ -99,6 +92,22 @@ class PreferentialStudy: else: raise RuntimeError("Unexpected trial type") storage.set_trial_system_attr(trial_id, _SYSTEM_ATTR_COMPARISON_READY, True) + + +def get_best_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]: + ready_trials = [ + t + for t in storage.get_all_trials( + study_id, + deepcopy=False, + states=(TrialState.COMPLETE, TrialState.RUNNING), + ) + if t.system_attrs.get(_SYSTEM_ATTR_COMPARISON_READY) is True + ] + preferences = get_preferences(study_id, storage) + worse_numbers = {worse for _, worse in preferences} + return [copy.deepcopy(t) for t in ready_trials if t.number not in worse_numbers] + def create_study( diff --git a/optuna_dashboard/preferential/_system_attrs.py b/optuna_dashboard/preferential/_system_attrs.py index 70567655..214d7466 100644 --- a/optuna_dashboard/preferential/_system_attrs.py +++ b/optuna_dashboard/preferential/_system_attrs.py @@ -5,42 +5,42 @@ import uuid import optuna from optuna.trial import FrozenTrial from optuna.trial import TrialState - +from optuna.storages import BaseStorage +from .._storage import get_study_summary _SYSTEM_ATTR_PREFIX_PREFERENCE = "preference:values" def report_preferences( - study: optuna.Study, - preferences: list[tuple[FrozenTrial, FrozenTrial]], + study_id: int, + storage: BaseStorage, + preferences: list[tuple[int, int]], # element is number of trail ) -> None: key = _SYSTEM_ATTR_PREFIX_PREFERENCE + str(uuid.uuid4()) - study._storage.set_study_system_attr( - study_id=study._study_id, + storage.set_study_system_attr( + study_id=study_id, key=key, - value=[(better.number, worse.number) for better, worse in preferences], + value=preferences, ) - - values = [0 for _ in study.directions] + trials = storage.get_all_trials(study_id, deepcopy=False) + directions = storage.get_study_directions(study_id) + values = [0 for _ in directions] for better, worse in preferences: - for t in (better, worse): - study.tell( - t.number, - values=values, - state=TrialState.COMPLETE, - skip_if_finished=True, - ) + for number in (better, worse): + trial_id = trials[number]._trial_id + if storage.check_trial_is_updatable(trial_id, trials[number].state): + storage.set_trial_state_values(trial_id, TrialState.COMPLETE, values) def get_preferences( - study: optuna.Study, - *, - deepcopy: bool = True, -) -> list[tuple[FrozenTrial, FrozenTrial]]: + study_id: int, + storage: BaseStorage, +) -> list[tuple[int, int]]: preferences: list[tuple[int, int]] = [] - for k, v in study.system_attrs.items(): + summary = get_study_summary(storage, study_id) + system_attrs = getattr(summary, "system_attrs", {}) + for k, v in system_attrs.items(): if not k.startswith(_SYSTEM_ATTR_PREFIX_PREFERENCE): continue preferences.extend(v) # type: ignore - trials = study.get_trials(deepcopy=deepcopy) - return [(trials[better], trials[worse]) for (better, worse) in preferences] + return preferences