mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-09-09 11:28:14 +08:00
74 lines
2.1 KiB
Python
74 lines
2.1 KiB
Python
import os
|
|
import tempfile
|
|
|
|
import optuna
|
|
from playwright.sync_api import Page
|
|
import pytest
|
|
|
|
from ..test_server import make_standalone_server
|
|
from ..utils import count_components
|
|
|
|
|
|
@pytest.fixture
|
|
def server_url(request: pytest.FixtureRequest) -> str:
|
|
return make_standalone_server(request)
|
|
|
|
|
|
def test_home(
|
|
page: Page,
|
|
server_url: str,
|
|
) -> None:
|
|
url = f"{server_url}"
|
|
page.goto(url)
|
|
element = page.get_by_role("heading")
|
|
assert element is not None
|
|
title = element.text_content()
|
|
assert title is not None
|
|
assert title == "Optuna Dashboard (Wasm ver.)"
|
|
|
|
|
|
def create_storage_file(filename: str, study_name: str, storage_type: str):
|
|
if storage_type == "rdb":
|
|
storage = optuna.storages.RDBStorage(f"sqlite:///{filename}")
|
|
elif storage_type == "journal":
|
|
storage = optuna.storages.JournalStorage(
|
|
optuna.storages.JournalFileStorage(f"{filename}"),
|
|
)
|
|
else:
|
|
assert False, f"Got an unexpected storage_type={storage_type}."
|
|
|
|
study = optuna.create_study(study_name=study_name, storage=storage)
|
|
|
|
def objective(trial: optuna.Trial) -> float:
|
|
x1 = trial.suggest_float("x1", 0, 10)
|
|
x2 = trial.suggest_float("x2", 0, 10)
|
|
return (x1 - 2) ** 2 + (x2 - 5) ** 2
|
|
|
|
study.optimize(objective, n_trials=100)
|
|
|
|
|
|
@pytest.mark.parametrize("storage_type", ["rdb", "journal"])
|
|
def test_load_storage(
|
|
page: Page,
|
|
server_url: str,
|
|
storage_type: str,
|
|
) -> None:
|
|
study_name = "single-objective"
|
|
url = f"{server_url}"
|
|
|
|
with tempfile.TemporaryDirectory() as dir:
|
|
with tempfile.NamedTemporaryFile() as fp:
|
|
filename = fp.name
|
|
path = os.path.join(dir, filename)
|
|
create_storage_file(filename, study_name, storage_type)
|
|
page.goto(url)
|
|
with page.expect_file_chooser() as fc_info:
|
|
page.get_by_role("button").filter(has_text="Storage").click()
|
|
file_chooser = fc_info.value
|
|
file_chooser.set_files(path)
|
|
|
|
page.get_by_role("link", name=study_name).click()
|
|
|
|
count = count_components(page, "MuiCard-root")
|
|
assert count == 4
|