From bb9a42b843fbcd1a0a4f57ec492af660f1e03c77 Mon Sep 17 00:00:00 2001 From: Contramundum Date: Mon, 14 Aug 2023 11:08:37 +0900 Subject: [PATCH] Add examples --- .../preferential-optimization/evaluator.py | 93 +++++++++++++++++++ .../preferential-optimization/generator.py | 63 +++++++++++++ 2 files changed, 156 insertions(+) create mode 100644 examples/preferential-optimization/evaluator.py create mode 100644 examples/preferential-optimization/generator.py diff --git a/examples/preferential-optimization/evaluator.py b/examples/preferential-optimization/evaluator.py new file mode 100644 index 00000000..c03f598b --- /dev/null +++ b/examples/preferential-optimization/evaluator.py @@ -0,0 +1,93 @@ +from __future__ import annotations + +import os +import shutil +import tempfile +import time +from typing import Callable +from typing import NoReturn +import uuid + +from optuna_dashboard.artifact.file_system import FileSystemBackend +import streamlit as st + +from optuna_dashboard.preferential import load_study + + +STORAGE_URL = "sqlite:///st-example.db" +artifact_path = os.path.join(os.path.dirname(__file__), "artifact") +artifact_backend = FileSystemBackend(base_path=artifact_path) +os.makedirs(artifact_path, exist_ok=True) + +n_comparison = 5 + + +def get_tmp_dir() -> str: + if "tmp_dir" not in st.session_state: + tmp_dir_name = str(uuid.uuid4()) + tmp_dir_path = os.path.join(tempfile.gettempdir(), tmp_dir_name) + os.makedirs(tmp_dir_path, exist_ok=True) + st.session_state.tmp_dir = tmp_dir_path + + return st.session_state.tmp_dir + + +def main() -> NoReturn: + tmpdir = get_tmp_dir() + study = load_study( + study_name="Preferential Optimization", + storage=STORAGE_URL, + ) + + # 1. 比較対象のTrialを取得 + comparison_trials = study.best_trials + + st.text("Which is the worst?") + + # 2. 各TrialのArtifact画像を並べて表示 + cols = st.columns(n_comparison) + finished_dict = {t.number: t for t in comparison_trials} + + col_is: dict[int, int] = st.session_state.get("col_is") + if col_is is None: + col_is = {} + col_is = {tn: col_i for (tn, col_i) in col_is.items() if tn in finished_dict} + + unoccupied_col_is = [i for i in range(len(cols)) if i not in col_is.values()] + for tn, col_i in zip([tn for tn in finished_dict if tn not in col_is], unoccupied_col_is): + col_is[tn] = col_i + st.session_state["col_is"] = col_is + + def on_click_factory(trial_number: int) -> Callable[[], None]: + def on_click() -> None: + better_trials = [t for t in comparison_trials if t.number != trial_number] + worse_trial = finished_dict[trial_number] + study.report_preference(better_trials, worse_trial) + + return on_click + + for trial_number, col_i in col_is.items(): + trial = finished_dict[trial_number] + col = cols[col_i] + + rgb_artifact_id = trial.user_attrs.get("rgb_artifact_id") + image_caption = trial.user_attrs.get("image_caption") + with col: + with artifact_backend.open(rgb_artifact_id) as fsrc: + tmp_img_path = os.path.join(tmpdir, rgb_artifact_id + ".png") + with open(tmp_img_path, "wb") as fdst: + shutil.copyfileobj(fsrc, fdst) + st.image(tmp_img_path, caption=image_caption) + st.button(str(trial_number), key=trial.number, on_click=on_click_factory(trial_number)) + + for i, col in enumerate(st.columns(n_comparison)): + if i >= len(comparison_trials): + continue + + if len(comparison_trials) < n_comparison: + time.sleep(0.1) + st.experimental_rerun() + + +if __name__ == "__main__": + main() diff --git a/examples/preferential-optimization/generator.py b/examples/preferential-optimization/generator.py new file mode 100644 index 00000000..d68f6c69 --- /dev/null +++ b/examples/preferential-optimization/generator.py @@ -0,0 +1,63 @@ +from __future__ import annotations + +import os +import tempfile +import time +from time import sleep +from typing import NoReturn + +import optuna +from optuna_dashboard.artifact import upload_artifact +from optuna_dashboard.artifact.file_system import FileSystemBackend +from PIL import Image + +from optuna_dashboard.preferential import create_study +from optuna_dashboard.preferential.samplers._gp import PreferentialGPSampler + + +STORAGE_URL = "sqlite:///st-example.db" +artifact_path = os.path.join(os.path.dirname(__file__), "artifact") +artifact_backend = FileSystemBackend(base_path=artifact_path) +os.makedirs(artifact_path, exist_ok=True) + +n_comparison = 5 + + +def main() -> NoReturn: + study = create_study( + study_name="Preferential Optimization", + storage=STORAGE_URL, + sampler=PreferentialGPSampler(), + load_if_exists=True, + ) + + with tempfile.TemporaryDirectory() as tmpdir: + while True: + if len(study.best_trials) >= n_comparison: + time.sleep(0.1) # Avoid busy-loop + continue + + trial = study.ask() + # 1. Ask new parameters + r = trial.suggest_int("r", 0, 255) + g = trial.suggest_int("g", 0, 255) + b = trial.suggest_int("b", 0, 255) + + # 2. Generate image + image_path = os.path.join(tmpdir, f"sample-{trial.number}.png") + image = Image.new("RGB", (320, 240), color=(r, g, b)) + # sleep(2.0) + image.save(image_path) + + # 3. Upload Artifact + artifact_id = upload_artifact(artifact_backend, trial, image_path) + trial.set_user_attr("rgb_artifact_id", artifact_id) + trial.set_user_attr("image_caption", f"(R, G, B) = ({r}, {g}, {b})") + print("RGB:", (r, g, b)) + + # 4. Mark comparison ready + study.mark_comparison_ready(trial) + + +if __name__ == "__main__": + main()