Add should_generate

This commit is contained in:
Contramundum
2023-08-31 17:31:12 +09:00
parent 4f6c78b18f
commit c05d0bd866
3 changed files with 48 additions and 2 deletions
@@ -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
+34 -1
View File
@@ -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,
)