mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-09-11 12:30:25 +08:00
Add should_generate
This commit is contained in:
@@ -35,7 +35,7 @@ def main() -> NoReturn:
|
||||
while True:
|
||||
# If n_comparison "best" trials (that are not reported bad) exists,
|
||||
# the generator waits for human evaluation.
|
||||
if len(study.best_trials) >= n_comparison:
|
||||
if study.should_generate():
|
||||
time.sleep(0.1) # Avoid busy-loop
|
||||
continue
|
||||
|
||||
|
||||
@@ -14,7 +14,7 @@ from optuna.trial import FrozenTrial
|
||||
from optuna.trial import TrialState
|
||||
from optuna_dashboard.preferential._system_attrs import get_preferences
|
||||
from optuna_dashboard.preferential._system_attrs import is_skipped_trial
|
||||
from optuna_dashboard.preferential._system_attrs import report_preferences
|
||||
from optuna_dashboard.preferential._system_attrs import report_preferences, get_n_generate, set_n_generate
|
||||
|
||||
|
||||
_logger = logging.get_logger(__name__)
|
||||
@@ -253,6 +253,36 @@ class PreferentialStudy:
|
||||
raise RuntimeError("Unexpected trial type")
|
||||
storage.set_trial_system_attr(trial_id, _SYSTEM_ATTR_COMPARISON_READY, True)
|
||||
|
||||
@property
|
||||
def n_generate(self) -> int:
|
||||
"""Return the number of trials that should be generated and shown to user.
|
||||
|
||||
:func:`~optuna_dashboard.preferential.PreferentialStudy.should_generate` returns
|
||||
:obj:`True` if the number of trials not reported bad and not skipped are less than
|
||||
:attr:`~optuna_dashboard.preferential.PreferentialStudy.n_generate`.
|
||||
"""
|
||||
system_attrs = self._study._storage.get_study_system_attrs(self._study._study_id)
|
||||
return get_n_generate(system_attrs)
|
||||
|
||||
def set_n_generate(self, n_generate: int) -> None:
|
||||
"""Set the number of trials that should be generated and shown to user.
|
||||
|
||||
:func:`~optuna_dashboard.preferential.PreferentialStudy.should_generate` returns
|
||||
:obj:`True` if the number of trials not reported bad and not skipped are less than
|
||||
:attr:`~optuna_dashboard.preferential.PreferentialStudy.n_generate`.
|
||||
"""
|
||||
return set_n_generate(self._study._study_id, self._study._storage, n_generate)
|
||||
|
||||
def should_generate(self) -> bool:
|
||||
"""Return whether the generator should generate a new trial now.
|
||||
|
||||
Returns :obj:`True` if the number of trials not reported bad and not skipped are less than
|
||||
:attr:`~optuna_dashboard.preferential.PreferentialStudy.n_generate`. Users are recommended
|
||||
to generate a new trial if this method returns :obj:`True`, and to wait for human evaluation
|
||||
if this method returns :obj:`False`.
|
||||
"""
|
||||
return len(self.best_trials) < self.n_generate
|
||||
|
||||
|
||||
def get_best_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]:
|
||||
preferences = get_preferences(study_id, storage)
|
||||
@@ -328,6 +358,9 @@ def create_study(
|
||||
study._storage.set_study_system_attr(
|
||||
study._study_id, _SYSTEM_ATTR_PREFERENTIAL_STUDY, True
|
||||
)
|
||||
study._storage.set_study_system_attr(
|
||||
study._study_id, _SYSTEM_ATTR_N_GENERATE, 4 # Default n_generate is 4
|
||||
)
|
||||
return PreferentialStudy(study)
|
||||
|
||||
except optuna.exceptions.DuplicatedStudyError:
|
||||
|
||||
@@ -9,6 +9,7 @@ from optuna.trial import TrialState
|
||||
|
||||
_SYSTEM_ATTR_PREFIX_PREFERENCE = "preference:values"
|
||||
_SYSTEM_ATTR_PREFIX_SKIP_TRIAL = "preference:skip_trial:"
|
||||
_SYSTEM_ATTR_N_GENERATE = "preference:n_generate"
|
||||
|
||||
|
||||
def report_preferences(
|
||||
@@ -60,3 +61,15 @@ def report_skip(
|
||||
def is_skipped_trial(trial_id: int, study_system_attrs: dict[str, Any]) -> bool:
|
||||
key = _SYSTEM_ATTR_PREFIX_SKIP_TRIAL + str(trial_id)
|
||||
return key in study_system_attrs
|
||||
|
||||
|
||||
def get_n_generate(study_system_attrs: dict[str, Any]) -> int:
|
||||
return study_system_attrs[_SYSTEM_ATTR_N_GENERATE]
|
||||
|
||||
def set_n_generate(study_id: int, n_generate: int, storage: BaseStorage) -> None:
|
||||
storage.set_study_system_attr(
|
||||
study_id=study_id,
|
||||
key=_SYSTEM_ATTR_N_GENERATE,
|
||||
value=n_generate,
|
||||
)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user