mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-09-09 11:28:14 +08:00
Update Human-in-the-loop example
This commit is contained in:
+9
-15
@@ -1,5 +1,6 @@
|
||||
import os
|
||||
import textwrap
|
||||
import time
|
||||
from typing import NoReturn
|
||||
|
||||
import optuna
|
||||
@@ -7,7 +8,6 @@ from optuna.trial import TrialState
|
||||
from optuna_dashboard import ChoiceWidget
|
||||
from optuna_dashboard import register_objective_form_widgets
|
||||
from optuna_dashboard import save_note
|
||||
from optuna_dashboard import set_objective_names
|
||||
from optuna_dashboard.artifact import get_artifact_path
|
||||
from optuna_dashboard.artifact import upload_artifact
|
||||
from optuna_dashboard.artifact.file_system import FileSystemBackend
|
||||
@@ -41,20 +41,17 @@ def suggest_and_generate_image(study: optuna.Study, artifact_backend: FileSystem
|
||||
save_note(trial, note)
|
||||
|
||||
|
||||
def start_optimization(
|
||||
storage: optuna.storages.BaseStorage, artifact_backend: FileSystemBackend
|
||||
) -> NoReturn:
|
||||
def start_optimization(artifact_backend: FileSystemBackend) -> NoReturn:
|
||||
# 1. Create Study
|
||||
sampler = optuna.samplers.TPESampler(constant_liar=True)
|
||||
study = optuna.create_study(
|
||||
study_name="Human-in-the-loop Optimization",
|
||||
storage=storage,
|
||||
sampler=sampler,
|
||||
storage="sqlite:///db.sqlite3",
|
||||
sampler=optuna.samplers.TPESampler(constant_liar=True, n_startup_trials=5),
|
||||
load_if_exists=True,
|
||||
)
|
||||
|
||||
# 2. Set an objective name
|
||||
set_objective_names(study, ["Looks like sunset color?"])
|
||||
study.set_metric_names(["Looks like sunset color?"])
|
||||
|
||||
# 3. Register ChoiceWidget
|
||||
register_objective_form_widgets(
|
||||
@@ -73,6 +70,7 @@ def start_optimization(
|
||||
while True:
|
||||
running_trials = study.get_trials(deepcopy=False, states=(TrialState.RUNNING,))
|
||||
if len(running_trials) >= n_batch:
|
||||
time.sleep(1) # Avoid busy-loop
|
||||
continue
|
||||
suggest_and_generate_image(study, artifact_backend)
|
||||
|
||||
@@ -80,11 +78,7 @@ def start_optimization(
|
||||
def main() -> NoReturn:
|
||||
tmp_path = os.path.join(os.path.dirname(__file__), "tmp")
|
||||
|
||||
# 1. Create RDBStorage
|
||||
url = "sqlite:///db.sqlite3"
|
||||
storage = optuna.storages.RDBStorage(url=url)
|
||||
|
||||
# 2. Create Artifact Storage
|
||||
# 1. Create Artifact Store
|
||||
artifact_path = os.path.join(os.path.dirname(__file__), "artifact")
|
||||
artifact_backend = FileSystemBackend(base_path=artifact_path)
|
||||
|
||||
@@ -94,8 +88,8 @@ def main() -> NoReturn:
|
||||
if not os.path.exists(tmp_path):
|
||||
os.mkdir(tmp_path)
|
||||
|
||||
# 3. Run optimize loop
|
||||
start_optimization(storage, artifact_backend)
|
||||
# 2. Run optimize loop
|
||||
start_optimization(artifact_backend)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
|
||||
Reference in New Issue
Block a user