mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-09-13 12:50:51 +08:00
fix preferential function
This commit is contained in:
committed by
moririn2528
parent
09676f2363
commit
e21159d40e
@@ -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(
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user