fix preferential function

This commit is contained in:
i23_moririn2528
2023-08-16 13:20:30 +09:00
committed by moririn2528
parent 09676f2363
commit e21159d40e
2 changed files with 43 additions and 34 deletions
+21 -12
View File
@@ -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(
+22 -22
View File
@@ -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