Merge remote-tracking branch 'origin/main' into support-3dmodel-artifacts

This commit is contained in:
hrntsm
2023-08-11 17:35:26 +09:00
36 changed files with 1199 additions and 371 deletions
+3
View File
@@ -0,0 +1,3 @@
[run]
concurrency = multiprocessing,thread
source = optuna_dashboard/
+44
View File
@@ -0,0 +1,44 @@
name: python coverage
on:
push:
branches:
- master
pull_request: {}
jobs:
coverage:
runs-on: ubuntu-latest
steps:
- name: Checkout
uses: actions/checkout@v3
- name: Setup Python
uses: actions/setup-python@v4
with:
python-version: '3.10'
architecture: x64
- name: Install dependencies
run: |
python -m pip install --progress-bar off --upgrade pip setuptools
pip install --progress-bar off .[optional]
pip install --progress-bar off .[test]
pip install --progress-bar off "optuna>=3.0.0"
pip install --progress-bar off .
echo 'import coverage; coverage.process_startup()' > sitecustomize.py
- name: Tests
env:
PYTHONPATH: . # To invoke sitecutomize.py
COVERAGE_PROCESS_START: .coveragerc # https://coverage.readthedocs.io/en/6.4.1/subprocess.html
COVERAGE_COVERAGE: yes # https://github.com/nedbat/coveragepy/blob/65bf33fc03209ffb01bbbc0d900017614645ee7a/coverage/control.py#L255-L261
run: |
coverage run --source=optuna_dashboard -m pytest python_tests
coverage combine
coverage xml
- name: Upload coverage to Codecov
uses: codecov/codecov-action@v3
with:
file: ./coverage.xml
+4 -2
View File
@@ -45,7 +45,8 @@ jobs:
# python_tests requires optuna>=3.0.0 since it imports FloatDistribution
run: |
python -m pip install --progress-bar off --upgrade pip setuptools
pip install streamlit boto3 moto[s3] pytest
pip install --progress-bar off .[optional]
pip install --progress-bar off .[test]
pip install --progress-bar off "optuna>=3.0.0"
pip install --progress-bar off .
- run: pytest python_tests
@@ -61,7 +62,8 @@ jobs:
- name: Install dependencies
run: |
python -m pip install --progress-bar off --upgrade pip setuptools
pip install streamlit boto3 moto[s3] pytest
pip install --progress-bar off .[optional]
pip install --progress-bar off .[test]
pip install --progress-bar off .
python -m pip install --progress-bar off --upgrade git+https://github.com/optuna/optuna.git
- run: pytest python_tests
+5
View File
@@ -28,6 +28,11 @@ docs/_generated/
rustlib/target/
rustlib/pkg/
# Test
.coverage
.coverage.*
coverage.xml
# Others
.envrc
.idea/
+43
View File
@@ -37,3 +37,46 @@ Please use `study.set_metric_names() <https://optuna.readthedocs.io/en/latest/re
* - ``optuna_dashboard.set_objective_names(study, ["objective 1", "objective 2"])``
- ``study.set_metric_names(["objective 1", "objective 2"])``
``upload_artifact()`` is deprecated. Please use ``optuna.artifacts.upload_artifact()`` instead.
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
:func:`~optuna_dashboard.artifact.upload_artifact` function has been ported to Optuna.
Please use `optuna.artifacts.upload_artifact <https://optuna.readthedocs.io/en/latest/reference/generated/optuna.artifacts.upload_artifact.html>`_ function instead.
.. list-table::
* - Deprecated APIs
- Corresponding Active APIs
* - ``optuna_dashboard.artifact.upload_artifact(artifact_backend, trial, fiel_path)``
- ``optuna.artifacts.upload_artifact(trial, file_path, artifact_store)``
Please note that the order of arguments is different between the deprecated and active APIs.
``FileSystemBackend`` is deprecated. Please use ``FileSystemArtifactStore`` instead.
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
:class:`~optuna_dashboard.artifact.file_system.FileSystemBackend` class has been ported to Optuna.
Please use `FileSystemArtifactStore <https://optuna.readthedocs.io/en/latest/reference/generated/optuna.artifacts.FileSystemArtifactStore.html>`_ class instead.
.. list-table::
* - Deprecated APIs
- Corresponding Active APIs
* - ``optuna_dashboard.artifact.file_system.FileSystemBackend(base_path)``
- ``optuna.artifacts.FileSystemArtifactStore(base_path)``
``Boto3Backend``` is deprecated. Please use ``Boto3ArtifactStore`` instead.
~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~
:class:`~optuna_dashboard.artifact.boto3.Boto3Backend` class has been ported to Optuna.
Please use `Boto3ArtifactStore <https://optuna.readthedocs.io/en/latest/reference/generated/optuna.artifacts.Boto3ArtifactStore.html>`_ class instead.
.. list-table::
* - Deprecated APIs
- Corresponding Active APIs
* - ``optuna_dashboard.artifact.boto3.Boto3Backend(bucket_name, client=None)``
- ``optuna.artifacts.Boto3ArtifactStore(bucket_name, client=None)``
+1 -1
View File
@@ -33,7 +33,7 @@ def start_streamlit() -> None:
study = optuna.load_study(
storage="sqlite:///streamlit-db.sqlite3", study_name="Human-in-the-loop Optimization"
)
selected_trial = st.sidebar.selectbox("一覧", study.trials, format_func=lambda t: t.number)
selected_trial = st.sidebar.selectbox("Trial", study.trials, format_func=lambda t: t.number)
if selected_trial is None:
return
+1 -1
View File
@@ -15,4 +15,4 @@ from ._note import get_note # noqa
from ._note import save_note # noqa
__version__ = "0.10.3"
__version__ = "0.12.0"
+38 -10
View File
@@ -7,6 +7,7 @@ import typing
from typing import Any
from typing import Optional
from typing import Union
import warnings
from bottle import Bottle
from bottle import redirect
@@ -36,10 +37,12 @@ from ._storage import get_trials
from ._storage_url import get_storage
from .artifact._backend import delete_all_artifacts
from .artifact._backend import register_artifact_route
from .artifact._backend_to_store import to_artifact_store
if typing.TYPE_CHECKING:
from _typeshed.wsgi import WSGIApplication
from optuna.artifacts._protocol import ArtifactStore
from optuna_dashboard.artifact.protocol import ArtifactBackend
@@ -54,7 +57,7 @@ cached_path_exists = functools.lru_cache(maxsize=10)(os.path.exists)
def create_app(
storage: BaseStorage,
artifact_backend: Optional[ArtifactBackend] = None,
artifact_store: Optional[ArtifactStore] = None,
debug: bool = False,
) -> Bottle:
app = Bottle()
@@ -76,7 +79,7 @@ def create_app(
@json_api_view
def api_meta() -> dict[str, Any]:
return {
"artifact_is_available": artifact_backend is not None,
"artifact_is_available": artifact_store is not None,
}
@app.get("/api/studies")
@@ -156,9 +159,8 @@ def create_app(
@app.delete("/api/studies/<study_id:int>")
@json_api_view
def delete_study(study_id: int) -> dict[str, Any]:
if artifact_backend is not None:
system_attrs = storage.get_study_system_attrs(study_id)
delete_all_artifacts(artifact_backend, system_attrs)
if artifact_store is not None:
delete_all_artifacts(artifact_store, storage, study_id)
try:
storage.delete_study(study_id)
@@ -347,7 +349,7 @@ def create_app(
return static_file(filename, root=STATIC_DIR)
register_rdb_migration_route(app, storage)
register_artifact_route(app, storage, artifact_backend)
register_artifact_route(app, storage, artifact_store)
return app
@@ -355,6 +357,8 @@ def run_server(
storage: Union[str, BaseStorage],
host: str = "localhost",
port: int = 8080,
artifact_store: Optional[ArtifactStore | ArtifactBackend] = None,
*,
artifact_backend: Optional[ArtifactBackend] = None,
) -> None:
"""Start running optuna-dashboard and blocks until the server terminates.
@@ -362,18 +366,42 @@ def run_server(
This function uses wsgiref module which is not intended for the production
use. If you want to run optuna-dashboard more secure and/or more fast,
please use WSGI server like Gunicorn or uWSGI via :func:`wsgi` function.
"""
app = create_app(get_storage(storage), artifact_backend=artifact_backend)
# TODO(c-bata): Remove artifact_backend keyword argument in the future release.
store: ArtifactStore | None = None
if artifact_store is not None:
store = to_artifact_store(artifact_store)
elif artifact_backend is not None:
warnings.warn(
"The `artifact_backend` argument is deprecated. "
"Please use `artifact_store` instead.",
DeprecationWarning,
)
store = to_artifact_store(artifact_backend)
app = create_app(get_storage(storage), artifact_store=store)
run(app, host=host, port=port)
def wsgi(
storage: Union[str, BaseStorage],
artifact_store: Optional[ArtifactBackend | ArtifactStore] = None,
*,
artifact_backend: Optional[ArtifactBackend] = None,
) -> WSGIApplication:
"""This function exposes WSGI interface for people who want to run on the
production-class WSGI servers like Gunicorn or uWSGI.
"""
return create_app(get_storage(storage), artifact_backend=artifact_backend)
# TODO(c-bata): Remove artifact_backend keyword argument in the future release.
store: ArtifactStore | None = None
if artifact_store is not None:
store = to_artifact_store(artifact_store)
elif artifact_backend is not None:
warnings.warn(
"The `artifact_backend` argument is deprecated. "
"Please use `artifact_store` instead.",
DeprecationWarning,
)
store = to_artifact_store(artifact_backend)
return create_app(get_storage(storage), artifact_store=store)
+15 -3
View File
@@ -12,17 +12,22 @@ from bottle import Bottle
from bottle import run
from optuna.storages import BaseStorage
from optuna.storages import RDBStorage
from optuna.version import __version__ as optuna_ver
from packaging import version
from . import __version__
from ._app import create_app
from ._sql_profiler import register_profiler_view
from ._storage_url import get_storage
from .artifact._backend_to_store import ArtifactBackendToStore
from .artifact.file_system import FileSystemBackend
if TYPE_CHECKING:
from typing import Literal
from optuna.artifacts._protocol import ArtifactStore
DEBUG = os.environ.get("OPTUNA_DASHBOARD_DEBUG") == "1"
SERVER_CHOICES = ["auto", "wsgiref", "gunicorn"]
@@ -113,10 +118,17 @@ def main() -> None:
storage: BaseStorage
storage = get_storage(args.storage, storage_class=args.storage_class)
artifact_backend = None
if args.artifact_dir is not None:
artifact_store: ArtifactStore | None
if args.artifact_dir is None:
artifact_store = None
elif version.parse(optuna_ver) >= version.Version("3.3.0"):
from optuna.artifacts import FileSystemArtifactStore
artifact_store = FileSystemArtifactStore(args.artifact_dir)
else:
artifact_backend = FileSystemBackend(args.artifact_dir)
app = create_app(storage, artifact_backend=artifact_backend, debug=DEBUG)
artifact_store = ArtifactBackendToStore(artifact_backend)
app = create_app(storage, artifact_store=artifact_store, debug=DEBUG)
if DEBUG and isinstance(storage, RDBStorage):
app = register_profiler_view(app, storage)
+1 -1
View File
@@ -187,7 +187,7 @@ def serialize_frozen_trial(
{k: trial_system_attrs[k] for k in trial_system_attrs if not k.startswith("dashboard")}
),
"note": note.get_note_from_system_attrs(study_system_attrs, trial._trial_id),
"artifacts": list_trial_artifacts(study_system_attrs, trial._trial_id),
"artifacts": list_trial_artifacts(study_system_attrs, trial),
"constraints": trial_system_attrs.get(CONSTRAINTS_KEY, []),
}
+61 -23
View File
@@ -6,6 +6,7 @@ import mimetypes
import os.path
from typing import TYPE_CHECKING
import uuid
import warnings
from bottle import BaseRequest
from bottle import Bottle
@@ -13,6 +14,7 @@ from bottle import HTTPResponse
from bottle import request
from bottle import response
import optuna
from optuna.trial import FrozenTrial
from .._bottle_util import json_api_view
from .._bottle_util import parse_data_uri
@@ -23,6 +25,7 @@ if TYPE_CHECKING:
from typing import Optional
from typing import TypedDict
from optuna.artifacts._protocol import ArtifactStore
from optuna.storages import BaseStorage
from .protocol import ArtifactBackend
@@ -56,14 +59,14 @@ def get_artifact_path(
def register_artifact_route(
app: Bottle, storage: BaseStorage, artifact_backend: Optional[ArtifactBackend]
app: Bottle, storage: BaseStorage, artifact_store: Optional[ArtifactStore]
) -> None:
@app.get("/artifacts/<study_id:int>/<trial_id:int>/<artifact_id:re:[0-9a-fA-F-]+>")
def proxy_artifact(study_id: int, trial_id: int, artifact_id: str) -> HTTPResponse | bytes:
if artifact_backend is None:
if artifact_store is None:
response.status = 400 # Bad Request
return b"Cannot access to the artifacts."
artifact_dict = _get_artifact_meta(storage, study_id, trial_id, artifact_id)
artifact_dict = get_artifact_meta(storage, study_id, trial_id, artifact_id)
if artifact_dict is None:
response.status = 404
return b"Not Found"
@@ -72,13 +75,14 @@ def register_artifact_route(
if encoding:
headers["Content-Encodings"] = encoding
fp = artifact_backend.open(artifact_id)
fp = artifact_store.open_reader(artifact_id)
return HTTPResponse(fp, headers=headers)
@app.post("/api/artifacts/<study_id:int>/<trial_id:int>")
@json_api_view
def upload_artifact_api(study_id: int, trial_id: int) -> dict[str, Any]:
if artifact_backend is None:
# TODO(c-bata): Use optuna.artifacts.upload_artifact()
if artifact_store is None:
response.status = 400 # Bad Request
return {"reason": "Cannot access to the artifacts."}
file = request.json.get("file")
@@ -89,7 +93,7 @@ def register_artifact_route(
_, data = parse_data_uri(file)
filename = request.json.get("filename", "")
artifact_id = str(uuid.uuid4())
artifact_backend.write(artifact_id, io.BytesIO(data))
artifact_store.write(artifact_id, io.BytesIO(data))
mimetype, encoding = mimetypes.guess_type(filename)
artifact = {
@@ -102,18 +106,22 @@ def register_artifact_route(
storage.set_study_system_attr(study_id, attr_key, json.dumps(artifact))
response.status = 201
trial = storage.get_trial(trial_id)
if trial is None:
response.status = 400
return {"reason": "Invalid study_id or trial_id"}
return {
"artifact_id": artifact_id,
"artifacts": list_trial_artifacts(storage.get_study_system_attrs(study_id), trial_id),
"artifacts": list_trial_artifacts(storage.get_study_system_attrs(study_id), trial),
}
@app.delete("/api/artifacts/<study_id:int>/<trial_id:int>/<artifact_id:re:[0-9a-fA-F-]+>")
@json_api_view
def delete_artifact(study_id: int, trial_id: int, artifact_id: str) -> dict[str, Any]:
if artifact_backend is None:
if artifact_store is None:
response.status = 400 # Bad Request
return {"reason": "Cannot access to the artifacts."}
artifact_backend.remove(artifact_id)
artifact_store.remove(artifact_id)
attr_key = _artifact_prefix(trial_id) + artifact_id
storage.set_study_system_attr(study_id, attr_key, json.dumps(None))
@@ -131,6 +139,12 @@ def upload_artifact(
) -> str:
"""Upload an artifact (files), which is associated with the trial.
.. warning::
This function is deprecated. Please use `optuna.artifacts.upload_artifact
<https://optuna.readthedocs.io/en/latest/reference/generated/optuna.artifacts.
upload_artifact.html>`_ instead.
Example:
.. code-block:: python
@@ -146,6 +160,12 @@ def upload_artifact(
upload_artifact(artifact_backend, trial, file_path)
return ...
"""
warnings.warn(
"upload_artifact() is deprecated. Please use optuna.artifacts.upload_artifact() instead.\n"
"See https://optuna-dashboard.readthedocs.io/en/latest/errors.html for details",
DeprecationWarning,
)
filename = os.path.basename(file_path)
storage = trial.storage
trial_id = trial._trial_id
@@ -170,31 +190,49 @@ def _artifact_prefix(trial_id: int) -> str:
return ARTIFACTS_ATTR_PREFIX + f"{trial_id}:"
def _get_artifact_meta(
def get_artifact_meta(
storage: BaseStorage, study_id: int, trial_id: int, artifact_id: str
) -> Optional[ArtifactMeta]:
study_system_attr = storage.get_study_system_attrs(study_id)
attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id
artifact_meta = study_system_attr.get(attr_key)
if artifact_meta is None:
return None
return json.loads(artifact_meta)
if artifact_meta is not None:
return json.loads(artifact_meta)
# See https://github.com/optuna/optuna/blob/f827582a8/optuna/artifacts/_upload.py#L71
trial_system_attrs = storage.get_trial_system_attrs(trial_id)
value = trial_system_attrs.get("artifacts:" + artifact_id)
if value is not None:
return json.loads(value)
return None
def delete_all_artifacts(backend: ArtifactBackend, study_system_attrs: dict[str, Any]) -> None:
artifact_meta_list: list[ArtifactMeta] = [
json.loads(value)
for key, value in study_system_attrs.items()
if key.startswith(ARTIFACTS_ATTR_PREFIX)
]
for meta in artifact_meta_list:
def delete_all_artifacts(backend: ArtifactStore, storage: BaseStorage, study_id: int) -> None:
artifact_metas = []
study_system_attrs = storage.get_study_system_attrs(study_id)
for trial in storage.get_all_trials(study_id):
trial_artifacts = list_trial_artifacts(study_system_attrs, trial)
artifact_metas.extend(trial_artifacts)
for meta in artifact_metas:
backend.remove(meta["artifact_id"])
def list_trial_artifacts(study_system_attrs: dict[str, Any], trial_id: int) -> list[ArtifactMeta]:
artifact_metas = [
def list_trial_artifacts(
study_system_attrs: dict[str, Any], trial: FrozenTrial
) -> list[ArtifactMeta]:
dashboard_artifact_metas = [
json.loads(value)
for key, value in study_system_attrs.items()
if key.startswith(_artifact_prefix(trial_id))
if key.startswith(_artifact_prefix(trial._trial_id))
]
# See https://github.com/optuna/optuna/blob/f827582a8/optuna/artifacts/_upload.py#L16
optuna_artifact_metas = [
json.loads(value)
for key, value in trial.system_attrs.items()
if key.startswith("artifacts:")
]
artifact_metas = dashboard_artifact_metas + optuna_artifact_metas
return [a for a in artifact_metas if a is not None]
@@ -0,0 +1,39 @@
from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from typing import BinaryIO
from typing import TypeGuard
from optuna.artifacts._protocol import ArtifactStore
from .protocol import ArtifactBackend
def is_artifact_backend(store: ArtifactBackend | ArtifactStore) -> TypeGuard[ArtifactBackend]:
return getattr(store, "open_reader", None) is None
def to_artifact_store(store: ArtifactBackend | ArtifactStore) -> ArtifactStore:
if is_artifact_backend(store):
return ArtifactBackendToStore(store)
# mypy cannot infer the type of `store` here.
return store # type: ignore
class ArtifactBackendToStore:
"""Converts a Dashboard's ArtifactBackend to Optuna's ArtifactStore."""
def __init__(self, artifact_backend: ArtifactBackend) -> None:
self._backend = artifact_backend
def open_reader(self, artifact_id: str) -> BinaryIO:
return self._backend.open(artifact_id)
def write(self, artifact_id: str, content_body: BinaryIO) -> None:
self._backend.write(artifact_id, content_body)
def remove(self, artifact_id: str) -> None:
self._backend.remove(artifact_id)
+12
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import io
import shutil
from typing import TYPE_CHECKING
import warnings
import boto3
from botocore.exceptions import ClientError
@@ -19,6 +20,12 @@ if TYPE_CHECKING:
class Boto3Backend:
"""An artifact backend for S3.
.. warning::
This class is deprecated. Please use `optuna.artifacts.Boto3ArtifactStore
<https://optuna.readthedocs.io/en/latest/reference/generated/optuna.artifacts.
Boto3ArtifactStore.html>`_ instead.
Example:
.. code-block:: python
@@ -44,6 +51,11 @@ class Boto3Backend:
# may close the source file object.
# See https://github.com/boto/boto3/issues/929
self._avoid_buf_copy = avoid_buf_copy
warnings.warn(
"Boto3Backend is deprecated. Please use Boto3ArtifactStore instead.\n"
"See https://optuna-dashboard.readthedocs.io/en/latest/errors.html for details",
DeprecationWarning,
)
def open(self, artifact_id: str) -> BinaryIO:
try:
+12
View File
@@ -3,6 +3,7 @@ from __future__ import annotations
import os
import shutil
from typing import TYPE_CHECKING
import warnings
from optuna_dashboard.artifact.exceptions import ArtifactNotFound
@@ -14,6 +15,12 @@ if TYPE_CHECKING:
class FileSystemBackend:
"""An artifact backend for file systems.
.. warning::
This class is deprecated. Please use `optuna.artifacts.FileSystemArtifactStore
<https://optuna.readthedocs.io/en/latest/reference/generated/optuna.artifacts.
FileSystemArtifactStore.html>`_ instead.
Example:
.. code-block:: python
@@ -32,6 +39,11 @@ class FileSystemBackend:
def __init__(self, base_path: str) -> None:
self._base_path = base_path
warnings.warn(
"FileSystemBackend is deprecated. Please use FileSystemArtifactStore instead.\n"
"See https://optuna-dashboard.readthedocs.io/en/latest/errors.html for details",
DeprecationWarning,
)
def open(self, artifact_id: str) -> BinaryIO:
filepath = os.path.join(self._base_path, artifact_id)
+12
View File
@@ -0,0 +1,12 @@
from __future__ import annotations
from optuna_dashboard.preferential._study import create_study
from optuna_dashboard.preferential._study import load_study
from optuna_dashboard.preferential._study import PreferentialStudy
__all__ = [
"PreferentialStudy",
"create_study",
"load_study",
]
+152
View File
@@ -0,0 +1,152 @@
from __future__ import annotations
import copy
from typing import Any
from typing import Container
from typing import Iterable
import optuna
from optuna import logging
from optuna.distributions import BaseDistribution
from optuna.samplers import BaseSampler
from optuna.samplers import RandomSampler
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 report_preferences
_logger = logging.get_logger(__name__)
_SYSTEM_ATTR_PREFERENTIAL_STUDY = "preference:is_preferential"
_SYSTEM_ATTR_COMPARISON_READY = "preference:comparison_ready"
class PreferentialStudy:
def __init__(self, study: optuna.Study) -> None:
self._study = study
@property
def trials(self) -> list[FrozenTrial]:
return self._study.trials
@property
def best_trials(self) -> list[FrozenTrial]:
ready_trials = [
t
for t in self._study.get_trials(
deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING)
)
if t.system_attrs.get(_SYSTEM_ATTR_COMPARISON_READY) is True
]
preferences = get_preferences(self._study, deepcopy=False)
worse_numbers = {worse.number for _, worse in preferences}
return [copy.deepcopy(t) for t in ready_trials if t.number not in worse_numbers]
@property
def study_name(self) -> str:
return self._study.study_name
@property
def user_attrs(self) -> dict[str, Any]:
return self._study.user_attrs
@property
def preferences(self) -> list[tuple[FrozenTrial, FrozenTrial]]:
return self.get_preferences(deepcopy=True)
def get_trials(
self,
deepcopy: bool = True,
states: Container[optuna.trial.TrialState] | None = None,
) -> list[FrozenTrial]:
return self._study.get_trials(deepcopy, states)
def ask(self, fixed_distributions: dict[str, BaseDistribution] | None = None) -> optuna.Trial:
return self._study.ask(fixed_distributions)
def add_trial(self, trial: FrozenTrial) -> None:
self._study.add_trial(trial)
def add_trials(self, trials: Iterable[FrozenTrial]) -> None:
self._study.add_trials(trials)
def report_preference(
self,
better_trials: FrozenTrial | list[FrozenTrial],
worse_trials: FrozenTrial | list[FrozenTrial],
) -> None:
if not isinstance(better_trials, list):
better_trials = [better_trials]
if not isinstance(worse_trials, list):
worse_trials = [worse_trials]
report_preferences(self._study, [(b, w) for b in better_trials for w in worse_trials])
def get_preferences(self, *, deepcopy: bool = True) -> list[tuple[FrozenTrial, FrozenTrial]]:
return get_preferences(self._study, deepcopy=deepcopy)
def set_user_attr(self, key: str, value: Any) -> None:
self._study.set_user_attr(key, value)
def mark_comparison_ready(self, trial_or_number: optuna.Trial | int) -> None:
storage = self._study._storage
if isinstance(trial_or_number, optuna.Trial):
trial_id = trial_or_number._trial_id
elif isinstance(trial_or_number, int):
trial_id = storage.get_trial_id_from_study_id_trial_number(
self._study._study_id, trial_or_number
)
else:
raise RuntimeError("Unexpected trial type")
storage.set_trial_system_attr(trial_id, _SYSTEM_ATTR_COMPARISON_READY, True)
def create_study(
*,
storage: str | optuna.storages.BaseStorage | None = None,
sampler: BaseSampler | None = None,
study_name: str | None = None,
load_if_exists: bool = False,
) -> PreferentialStudy:
try:
study = optuna.create_study(
storage=storage,
sampler=sampler or RandomSampler(),
study_name=study_name,
)
study._storage.set_study_system_attr(
study._study_id, _SYSTEM_ATTR_PREFERENTIAL_STUDY, True
)
return PreferentialStudy(study)
except optuna.exceptions.DuplicatedStudyError:
if load_if_exists:
assert study_name is not None
assert storage is not None
_logger.info(
"Using an existing study with name '{}' instead of "
"creating a new one.".format(study_name)
)
return load_study(
study_name=study_name,
storage=storage,
sampler=sampler,
)
else:
raise
def load_study(
*,
study_name: str | None,
storage: str | optuna.storages.BaseStorage,
sampler: BaseSampler | None = None,
) -> PreferentialStudy:
study = optuna.load_study(
study_name=study_name, storage=storage, sampler=sampler or RandomSampler()
)
system_attrs = study._storage.get_study_system_attrs(study._study_id)
if not system_attrs.get(_SYSTEM_ATTR_PREFERENTIAL_STUDY):
raise ValueError("The study is not a PreferentialStudy.")
return PreferentialStudy(study)
@@ -0,0 +1,46 @@
from __future__ import annotations
import uuid
import optuna
from optuna.trial import FrozenTrial
from optuna.trial import TrialState
_SYSTEM_ATTR_PREFIX_PREFERENCE = "preference:values"
def report_preferences(
study: optuna.Study,
preferences: list[tuple[FrozenTrial, FrozenTrial]],
) -> None:
key = _SYSTEM_ATTR_PREFIX_PREFERENCE + str(uuid.uuid4())
study._storage.set_study_system_attr(
study_id=study._study_id,
key=key,
value=[(better.number, worse.number) for better, worse in preferences],
)
values = [0 for _ in study.directions]
for better, worse in preferences:
for t in (better, worse):
study.tell(
t.number,
values=values,
state=TrialState.COMPLETE,
skip_if_finished=True,
)
def get_preferences(
study: optuna.Study,
*,
deepcopy: bool = True,
) -> list[tuple[FrozenTrial, FrozenTrial]]:
preferences: list[tuple[int, int]] = []
for k, v in study.system_attrs.items():
if not k.startswith(_SYSTEM_ATTR_PREFIX_PREFERENCE):
continue
preferences.extend(v) # type: ignore
trials = study.get_trials(deepcopy=deepcopy)
return [(trials[better], trials[worse]) for (better, worse) in preferences]
+8 -2
View File
@@ -34,6 +34,7 @@ type LocalStorageReloadInterval = {
reloadInterval?: number
}
// eslint-disable-next-line @typescript-eslint/explicit-module-boundary-types
export const actionCreator = () => {
const { enqueueSnackbar } = useSnackbar()
const [studySummaries, setStudySummaries] =
@@ -432,8 +433,13 @@ export const actionCreator = () => {
const reader = new FileReader()
setUploading(true)
reader.readAsDataURL(file)
reader.onload = (upload: any) => {
uploadArtifactAPI(studyId, trialId, file.name, upload.target.result)
reader.onload = (upload: ProgressEvent<FileReader>) => {
uploadArtifactAPI(
studyId,
trialId,
file.name,
upload.target?.result as string
)
.then((res) => {
setUploading(false)
const index = studyDetails[studyId].trials.findIndex(
@@ -27,7 +27,7 @@ import { actionCreator } from "../action"
import { studySummariesState, studyDetailsState } from "../state"
import { AppDrawer } from "./AppDrawer"
import { GraphEdfMultiStudies } from "./GraphEdf"
import { GraphHistoryMultiStudies } from "./GraphHistory"
import { GraphHistory } from "./GraphHistory"
import { useNavigate, useLocation } from "react-router-dom"
const useQuery = (): URLSearchParams => {
@@ -199,7 +199,9 @@ export const CompareStudies: FC<{
alignItems: "flex-start",
}}
>
<ListItemText primary={study.study_name} />
<ListItemText
primary={`${study.study_id}. ${study.study_name}`}
/>
<Box
sx={{
display: "flex",
@@ -207,12 +209,6 @@ export const CompareStudies: FC<{
width: "100%",
}}
>
<Typography
color={theme.palette.grey.A400}
sx={{ p: theme.spacing(0, 1) }}
>
{`# ${study.study_id}`}
</Typography>
<Chip
color="primary"
label={
@@ -220,7 +216,15 @@ export const CompareStudies: FC<{
? `${study.directions.length} objective`
: `${study.directions.length} objectives`
}
sx={{ margin: theme.spacing(0) }}
size="small"
variant="outlined"
/>
<span style={{ margin: theme.spacing(0.5) }} />
<Chip
color="secondary"
label={study.directions
.map((d) => (d === "maximize" ? "max" : "min"))
.join(", ")}
size="small"
variant="outlined"
/>
@@ -313,7 +317,7 @@ const StudiesGraph: FC<{ studies: StudySummary[] }> = ({ studies }) => {
}}
>
<CardContent>
<GraphHistoryMultiStudies
<GraphHistory
studies={showStudyDetails}
includePruned={includePruned}
logScale={logScale}
+7 -4
View File
@@ -19,6 +19,9 @@ import { Clear } from "@mui/icons-material"
type Order = "asc" | "desc"
// eslint-disable-next-line @typescript-eslint/no-explicit-any
type Value = any
const defaultRowsPerPageOption = [10, 50, 100, { label: "All", value: -1 }]
interface DataGridColumn<T> {
@@ -33,7 +36,7 @@ interface DataGridColumn<T> {
interface RowFilter {
columnIdx: number
value: any
value: Value
}
function DataGrid<T>(props: {
@@ -45,7 +48,7 @@ function DataGrid<T>(props: {
initialRowsPerPage?: number
rowsPerPageOption?: Array<number | { value: number; label: string }>
defaultFilter?: (row: T) => boolean
}) {
}): React.ReactElement {
const { columns, rows, keyField, dense, collapseBody, defaultFilter } = props
let { initialRowsPerPage, rowsPerPageOption } = props
const [order, setOrder] = React.useState<Order>("asc")
@@ -81,7 +84,7 @@ function DataGrid<T>(props: {
const fieldAlreadyFiltered = (columnIdx: number): boolean =>
filters.some((f) => f.columnIdx === columnIdx)
const handleClickFilterCell = (columnIdx: number, value: any) => {
const handleClickFilterCell = (columnIdx: number, value: Value) => {
if (fieldAlreadyFiltered(columnIdx)) {
return
}
@@ -242,7 +245,7 @@ function DataGridRow<T>(props: {
row: T
keyField: keyof T
collapseBody?: (rowIndex: number) => React.ReactNode
handleClickFilterCell: (columnIdx: number, value: any) => void
handleClickFilterCell: (columnIdx: number, value: Value) => void
}) {
const {
columns,
+14 -281
View File
@@ -1,6 +1,7 @@
import * as plotly from "plotly.js-dist-min"
import React, { ChangeEvent, FC, useEffect, useState } from "react"
import {
Box,
Grid,
FormControl,
FormLabel,
@@ -16,10 +17,9 @@ import {
} from "@mui/material"
import { plotlyDarkTemplate } from "./PlotlyDarkMode"
import {
useFilteredTrials,
useFilteredTrialsFromStudies,
Target,
useObjectiveAndUserAttrTargets,
useObjectiveAndUserAttrTargetsFromStudies,
} from "../trialFilter"
const plotDomId = "graph-history"
@@ -32,141 +32,6 @@ interface HistoryPlotInfo {
}
export const GraphHistory: FC<{
study: StudyDetail | null
logScale: boolean
includePruned: boolean
}> = ({ study, logScale, includePruned }) => {
const theme = useTheme()
const [xAxis, setXAxis] = useState<
"number" | "datetime_start" | "datetime_complete"
>("number")
const [markerSize, setMarkerSize] = useState<number>(5)
const [targets, selected, setTarget] = useObjectiveAndUserAttrTargets(study)
const trials = useFilteredTrials(study, [selected], !includePruned)
useEffect(() => {
if (study !== null) {
plotHistory(
trials,
study.directions,
selected,
xAxis,
logScale,
theme.palette.mode,
study?.objective_names,
markerSize
)
}
}, [
trials,
study?.directions,
selected,
logScale,
xAxis,
theme.palette.mode,
study?.objective_names,
markerSize,
])
const handleObjectiveChange = (event: SelectChangeEvent<string>) => {
setTarget(event.target.value)
}
const handleXAxisChange = (e: ChangeEvent<HTMLInputElement>) => {
if (e.target.value === "number") {
setXAxis("number")
} else if (e.target.value === "datetime_start") {
setXAxis("datetime_start")
} else if (e.target.value === "datetime_complete") {
setXAxis("datetime_complete")
}
}
return (
<Grid container direction="row">
<Grid
item
xs={3}
container
direction="column"
sx={{ paddingRight: theme.spacing(2) }}
>
<Typography
variant="h6"
sx={{ margin: "1em 0", fontWeight: theme.typography.fontWeightBold }}
>
History
</Typography>
{targets.length >= 2 ? (
<FormControl
component="fieldset"
sx={{ marginBottom: theme.spacing(2) }}
>
<FormLabel component="legend">y Axis</FormLabel>
<Select
value={selected.identifier()}
onChange={handleObjectiveChange}
>
{targets.map((t, i) => (
<MenuItem value={t.identifier()} key={i}>
{t.toLabel(study?.objective_names)}
</MenuItem>
))}
</Select>
</FormControl>
) : null}
<FormControl
component="fieldset"
sx={{ marginBottom: theme.spacing(2) }}
>
<FormLabel component="legend">X-axis:</FormLabel>
<RadioGroup
aria-label="gender"
name="gender1"
value={xAxis}
onChange={handleXAxisChange}
>
<FormControlLabel
value="number"
control={<Radio />}
label="Number"
/>
<FormControlLabel
value="datetime_start"
control={<Radio />}
label="Datetime start"
/>
<FormControlLabel
value="datetime_complete"
control={<Radio />}
label="Datetime complete"
/>
</RadioGroup>
</FormControl>
<FormControl>
<FormLabel component="legend">Marker size:</FormLabel>
<Slider
defaultValue={5}
marks={true}
min={1}
max={10}
step={1}
onChange={(e) => {
// @ts-ignore
setMarkerSize(e.target.value as number)
}}
/>
</FormControl>
</Grid>
<Grid item xs={9}>
<div id={plotDomId} />
</Grid>
</Grid>
)
}
export const GraphHistoryMultiStudies: FC<{
studies: StudyDetail[]
logScale: boolean
includePruned: boolean
@@ -177,10 +42,8 @@ export const GraphHistoryMultiStudies: FC<{
>("number")
const [markerSize, setMarkerSize] = useState<number>(5)
// TODO(umezawa): Prepare targets with all studies.
const [targets, selected, setTarget] = useObjectiveAndUserAttrTargets(
studies.length !== 0 ? studies[0] : null
)
const [targets, selected, setTarget] =
useObjectiveAndUserAttrTargetsFromStudies(studies)
const trials = useFilteredTrialsFromStudies(
studies,
@@ -198,7 +61,7 @@ export const GraphHistoryMultiStudies: FC<{
})
useEffect(() => {
plotHistoryMultiStudies(
plotHistory(
historyPlotInfos,
selected,
xAxis,
@@ -206,7 +69,7 @@ export const GraphHistoryMultiStudies: FC<{
theme.palette.mode,
markerSize
)
}, [studies, selected, logScale, xAxis, theme.palette.mode])
}, [studies, selected, logScale, xAxis, theme.palette.mode, markerSize])
const handleObjectiveChange = (event: SelectChangeEvent<string>) => {
setTarget(event.target.value)
@@ -299,148 +162,18 @@ export const GraphHistoryMultiStudies: FC<{
</FormControl>
</Grid>
<Grid item xs={9}>
<div id={plotDomId} />
<Box
id={plotDomId}
sx={{
height: "450px",
}}
/>
</Grid>
</Grid>
)
}
const plotHistory = (
trials: Trial[],
directions: StudyDirection[],
target: Target,
xAxis: "number" | "datetime_start" | "datetime_complete",
logScale: boolean,
mode: string,
objectiveNames?: string[],
markerSize: number
) => {
if (document.getElementById(plotDomId) === null) {
return
}
const layout: Partial<plotly.Layout> = {
margin: {
l: 50,
t: 0,
r: 50,
b: 0,
},
yaxis: {
title: target.toLabel(objectiveNames),
type: logScale ? "log" : "linear",
},
xaxis: {
title: xAxis === "number" ? "Trial" : "Time",
type: xAxis === "number" ? "linear" : "date",
},
showlegend: true,
uirevision: "true",
template: mode === "dark" ? plotlyDarkTemplate : {},
}
if (trials.length === 0) {
plotly.react(plotDomId, [], layout)
return
}
const feasibleTrials: Trial[] = []
const infeasibleTrials: Trial[] = []
trials.forEach((t) => {
if (t.constraints.every((c) => c <= 0)) {
feasibleTrials.push(t)
} else {
infeasibleTrials.push(t)
}
})
const getAxisX = (trial: Trial): number | Date => {
return xAxis === "number"
? trial.number
: xAxis === "datetime_start"
? trial.datetime_start!
: trial.datetime_complete!
}
const plotData: Partial<plotly.PlotData>[] = [
{
x: feasibleTrials.map(getAxisX),
y: feasibleTrials.map(
(t: Trial): number => target.getTargetValue(t) as number
),
name: target.toLabel(objectiveNames),
marker: {
size: markerSize,
},
mode: "markers",
type: "scatter",
},
]
const objectiveId = target.getObjectiveId()
if (objectiveId !== null) {
const xForLinePlot: (number | Date)[] = []
const yForLinePlot: number[] = []
let currentBest: number | null = null
for (let i = 0; i < feasibleTrials.length; i++) {
const t = feasibleTrials[i]
if (currentBest === null) {
currentBest = t.values![objectiveId] as number
xForLinePlot.push(getAxisX(t))
yForLinePlot.push(t.values![objectiveId] as number)
} else if (
directions[objectiveId] === "maximize" &&
t.values![objectiveId] > currentBest
) {
const p = trials[i - 1]
if (!xForLinePlot.includes(getAxisX(p))) {
xForLinePlot.push(getAxisX(p))
yForLinePlot.push(currentBest)
}
currentBest = t.values![objectiveId] as number
xForLinePlot.push(getAxisX(t))
yForLinePlot.push(t.values![objectiveId] as number)
} else if (
directions[objectiveId] === "minimize" &&
t.values![objectiveId] < currentBest
) {
const p = feasibleTrials[i - 1]
if (!xForLinePlot.includes(getAxisX(p))) {
xForLinePlot.push(getAxisX(p))
yForLinePlot.push(currentBest)
}
currentBest = t.values![objectiveId] as number
xForLinePlot.push(getAxisX(t))
yForLinePlot.push(t.values![objectiveId] as number)
}
}
xForLinePlot.push(getAxisX(trials[trials.length - 1]))
yForLinePlot.push(yForLinePlot[yForLinePlot.length - 1])
plotData.push({
x: xForLinePlot,
y: yForLinePlot,
name: "Best Value",
mode: "lines",
type: "scatter",
})
}
plotData.push({
x: infeasibleTrials.map(getAxisX),
y: infeasibleTrials.map(
(t: Trial): number => target.getTargetValue(t) as number
),
name: "Infeasible Trial",
marker: {
size: markerSize,
color: mode === "dark" ? "#666666" : "#cccccc",
},
mode: "markers",
type: "scatter",
showlegend: false,
})
plotly.react(plotDomId, plotData, layout)
}
const plotHistoryMultiStudies = (
historyPlotInfos: HistoryPlotInfo[],
target: Target,
xAxis: "number" | "datetime_start" | "datetime_complete",
@@ -481,8 +214,8 @@ const plotHistoryMultiStudies = (
return xAxis === "number"
? trial.number
: xAxis === "datetime_start"
? trial.datetime_start!
: trial.datetime_complete!
? trial.datetime_start ?? new Date()
: trial.datetime_complete ?? new Date()
}
const plotData: Partial<plotly.PlotData>[] = []
@@ -109,12 +109,17 @@ const makeScatterObject = (
objectiveYId: number,
hovertemplate: string,
dominated: boolean,
feasible: boolean
feasible: boolean,
mode: string
): Partial<plotly.PlotData> => {
const marker = makeMarker(trials, dominated, feasible)
const marker = makeMarker(trials, dominated, feasible, mode)
return {
x: trials.map((t) => t.values![objectiveXId] as number),
y: trials.map((t) => t.values![objectiveYId] as number),
x: trials.map((t) =>
t.values ? (t.values[objectiveXId] as number) : null
),
y: trials.map((t) =>
t.values ? (t.values[objectiveYId] as number) : null
),
text: trials.map((t) => makeHovertext(t)),
mode: "markers",
hovertemplate: hovertemplate,
@@ -126,7 +131,8 @@ const makeScatterObject = (
const makeMarker = (
trials: Trial[],
dominated: boolean,
feasible: boolean
feasible: boolean,
mode: string
): Partial<plotly.PlotData> => {
if (feasible && dominated) {
return {
@@ -154,7 +160,7 @@ const makeMarker = (
} else {
return {
// @ts-ignore
color: "#cccccc",
color: mode === "dark" ? "#666666" : "#cccccc",
}
}
}
@@ -234,7 +240,8 @@ const plotParetoFront = (
? "%{text}<extra>Trial</extra>"
: "%{text}<extra>Feasible Trial</extra>",
true,
true
true,
mode
),
makeScatterObject(
feasibleTrials.filter((t, i) => !dominatedTrials[i]),
@@ -242,7 +249,8 @@ const plotParetoFront = (
objectiveYId,
"%{text}<extra>Best Trial</extra>",
false,
true
true,
mode
),
makeScatterObject(
infeasibleTrials,
@@ -250,7 +258,8 @@ const plotParetoFront = (
objectiveYId,
"%{text}<extra>Infeasible Trial</extra>",
false,
false
false,
mode
),
]
@@ -96,7 +96,7 @@ export const StudyHistory: FC<{ studyId: number }> = ({ studyId }) => {
>
<CardContent>
<GraphHistory
study={studyDetail}
studies={studyDetail !== null ? [studyDetail] : []}
includePruned={includePruned}
logScale={logScale}
/>
+24 -6
View File
@@ -1,4 +1,5 @@
import React, { FC, useEffect, useState } from "react"
import React, { FC, useEffect, useMemo, useState } from "react"
import { useNavigate, useLocation } from "react-router-dom"
import { useRecoilValue } from "recoil"
import { Link } from "react-router-dom"
import {
@@ -56,10 +57,20 @@ export const StudyList: FC<{
useDeleteStudyDialog()
const [openRenameStudyDialog, renderRenameStudyDialog] =
useRenameStudyDialog(studies)
const [sortBy, setSortBy] = useState<"id-asc" | "id-desc">("id-asc")
const navigate = useNavigate()
const useQuery = (): URLSearchParams => {
const { search } = useLocation()
return useMemo(() => new URLSearchParams(search), [search])
}
const query = useQuery()
const initialSortBy =
query.get("studies_order_by") === "desc" ? "desc" : "asc"
const [sortBy, setSortBy] = useState<"asc" | "desc">(initialSortBy)
let filteredStudies = studies.filter((s) => !studyFilter(s))
if (sortBy === "id-desc") {
if (sortBy === "desc") {
filteredStudies = filteredStudies.reverse()
}
@@ -67,6 +78,13 @@ export const StudyList: FC<{
action.updateStudySummaries()
}, [])
useEffect(() => {
query.set("studies_order_by", sortBy)
navigate(`${location.pathname}?${query.toString()}`, {
replace: true,
})
}, [sortBy])
const Select = styled(TextField)(({ theme }) => ({
"& .MuiInputBase-input": {
// vertical padding + font size from searchIcon
@@ -98,11 +116,11 @@ export const StudyList: FC<{
select
value={sortBy}
onChange={(e) => {
setSortBy(e.target.value as "id-asc" | "id-desc")
setSortBy(e.target.value as "asc" | "desc")
}}
>
<MenuItem value={"id-asc"}>Sort ascending</MenuItem>
<MenuItem value={"id-desc"}>Sort descending</MenuItem>
<MenuItem value={"asc"}>Sort ascending</MenuItem>
<MenuItem value={"desc"}>Sort descending</MenuItem>
</Select>
</Box>
)
+58
View File
@@ -192,3 +192,61 @@ export const useObjectiveAndUserAttrTargets = (
)
return [targetList, selectedTarget, setTargetIdent]
}
export const useObjectiveAndUserAttrTargetsFromStudies = (
studies: StudyDetail[]
): [Target[], Target, (ident: string) => void] => {
const defaultTarget = new Target("objective", 0)
const [selected, setTargetIdent] = useState<string>(
defaultTarget.identifier()
)
const minDirections = useMemo<number>(() => {
if (studies.length === 0) {
return 0
}
return studies.reduce((acc, study) => {
return Math.min(acc, study.directions.length)
}, Number.MAX_VALUE)
}, [studies])
const intersect = (arrays: AttributeSpec[][]) => {
const atrEqual = (obj1: AttributeSpec, obj2: AttributeSpec) => {
return obj1.key === obj2.key
}
return arrays.reduce((a, b) =>
a.filter((c) => b.some((d) => atrEqual(c, d)))
)
}
const attrTargets = useMemo<Target[]>(() => {
if (studies.length === 0) {
return []
}
const intersection = intersect(
studies.map((study) => study.union_user_attrs)
)
return intersection
.filter((attr) => attr.sortable)
.map((attr) => new Target("user_attr", attr.key))
}, [studies])
const targetList = useMemo<Target[]>(() => {
if (studies !== null) {
return [
...Array.from(
{ length: minDirections },
(_, i) => new Target("objective", i)
),
...attrTargets,
]
} else {
return [defaultTarget]
}
}, [minDirections, attrTargets])
const selectedTarget = useMemo<Target>(
() => targetList.find((t) => t.identifier() === selected) || defaultTarget,
[targetList, selected]
)
return [targetList, selectedTarget, setTargetIdent]
}
+1 -1
View File
@@ -7,7 +7,7 @@
"scripts": {
"fmt": "prettier --write \"{optuna_dashboard/ts,typescript_tests,standalone_app/src,vscode/src}/**/*.{ts,tsx}\"",
"lint": "npm run lint:eslint && npm run lint:fmt",
"lint:eslint": "eslint . --ext .ts,.tsx",
"lint:eslint": "eslint . --ext .ts,.tsx --max-warnings 0",
"lint:fmt": "prettier --list-different \"{optuna_dashboard/ts,typescript_tests,standalone_app/src,vscode/src}/**/*.{ts,tsx}\"",
"watch": "NODE_ENV=development TYPESCRIPT_LOADER=esbuild-loader webpack --watch",
"build": "webpack",
+12
View File
@@ -33,6 +33,18 @@ dependencies = [
]
dynamic = ["version"]
[project.optional-dependencies]
test = [
"coverage",
"pytest",
"moto[s3]",
]
optional = [
"streamlit",
"boto3",
]
[project.scripts]
optuna-dashboard = "optuna_dashboard._cli:main"
@@ -0,0 +1,83 @@
from __future__ import annotations
import os
import tempfile
import optuna
from optuna.version import __version__ as optuna_ver
from optuna_dashboard.artifact._backend import delete_all_artifacts
from optuna_dashboard.artifact._backend import get_artifact_meta
from optuna_dashboard.artifact._backend import list_trial_artifacts
from packaging import version
import pytest
@pytest.mark.skipif(
version.parse(optuna_ver) < version.Version("3.3.0"),
reason="Artifact is not implemented yet in Optuna",
)
def test_list_optuna_trial_artifacts() -> None:
from optuna.artifacts import FileSystemArtifactStore
from optuna.artifacts import upload_artifact
storage = optuna.storages.InMemoryStorage()
study = optuna.create_study(storage=storage)
dummy_content = b"dummy content"
with tempfile.TemporaryDirectory() as tmpdir:
artifact_store = FileSystemArtifactStore(tmpdir)
trial = study.ask()
with tempfile.NamedTemporaryFile() as f:
f.write(dummy_content)
f.flush()
upload_artifact(trial, f.name, artifact_store=artifact_store)
study.tell(trial, 0.0)
study_system_attrs = storage.get_study_system_attrs(study._study_id)
frozen_trial = storage.get_trial(trial._trial_id)
artifact_meta_list = list_trial_artifacts(study_system_attrs, frozen_trial)
assert len(artifact_meta_list) == 1
artifact_id = artifact_meta_list[0]["artifact_id"]
with artifact_store.open_reader(artifact_id) as reader:
assert reader.read() == dummy_content
artifact_meta = get_artifact_meta(
storage=storage,
study_id=study._study_id,
trial_id=trial._trial_id,
artifact_id=artifact_id,
)
assert artifact_meta is not None
@pytest.mark.skipif(
version.parse(optuna_ver) < version.Version("3.3.0"),
reason="Artifact is not implemented yet in Optuna",
)
def test_delete_optuna_study_artifacts() -> None:
from optuna.artifacts import FileSystemArtifactStore
from optuna.artifacts import upload_artifact
storage = optuna.storages.InMemoryStorage()
study = optuna.create_study(storage=storage)
with tempfile.TemporaryDirectory() as tmpdir:
dummy_file_path = os.path.join(tmpdir, "dummy.txt")
with open(dummy_file_path, "wb") as f:
f.write(b"dummy content")
f.flush()
artifact_store = FileSystemArtifactStore(tmpdir)
def objective(trial: optuna.Trial) -> float:
upload_artifact(trial, dummy_file_path, artifact_store=artifact_store)
return 0.0
study.optimize(objective, n_trials=10)
assert len(os.listdir(tmpdir)) == 11 # 10 artifacts + dummy.txt
delete_all_artifacts(artifact_store, storage, study._study_id)
assert len(os.listdir(tmpdir)) == 1 # dummy.txt only
+326
View File
@@ -0,0 +1,326 @@
from __future__ import annotations
import copy
import multiprocessing
import pickle
from typing import Callable
from unittest.mock import patch
import uuid
from optuna import copy_study
from optuna import create_trial
from optuna import delete_study
from optuna import distributions
from optuna import Trial
from optuna.exceptions import DuplicatedStudyError
from optuna.trial import TrialState
from optuna_dashboard.preferential import create_study
from optuna_dashboard.preferential import load_study
import pytest
from ..storage_supplier import parametrize_storages
from ..storage_supplier import StorageSupplier
@parametrize_storages
def test_study_set_and_get_user_attrs(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
study = create_study(storage=storage)
study.set_user_attr("dataset", "MNIST")
assert study.user_attrs["dataset"] == "MNIST"
@parametrize_storages
def test_report_and_get_preferences(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
study = create_study(storage=storage)
assert len(study.preferences) == 0
for _ in range(2):
trial = study.ask()
trial.suggest_float("x", 0, 1)
study.mark_comparison_ready(trial)
better, worse = study.trials
study.report_preference(better, worse)
assert len(study.preferences) == 1
actual_better, actual_worse = study.preferences[0]
assert actual_better.number == better.number
assert actual_worse.number == worse.number
def test_study_pickle() -> None:
study_1 = create_study()
for _ in range(10):
study_1.ask()
assert len(study_1.trials) == 10
dumped_bytes = pickle.dumps(study_1)
study_2 = pickle.loads(dumped_bytes)
assert len(study_2.trials) == 10
for _ in range(10):
study_2.ask()
assert len(study_2.trials) == 20
@parametrize_storages
def test_create_study(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
# Test creating a new study.
study = create_study(storage=storage, load_if_exists=False)
# Test `load_if_exists=True` with existing study.
create_study(study_name=study.study_name, storage=storage, load_if_exists=True)
with pytest.raises(DuplicatedStudyError):
create_study(study_name=study.study_name, storage=storage, load_if_exists=False)
@parametrize_storages
def test_load_study(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
if storage is None:
# `InMemoryStorage` can not be used with `load_study` function.
return
study_name = str(uuid.uuid4())
with pytest.raises(KeyError):
# Test loading an unexisting study.
load_study(study_name=study_name, storage=storage)
# Create a new study.
created_study = create_study(study_name=study_name, storage=storage)
# Test loading an existing study.
loaded_study = load_study(study_name=study_name, storage=storage)
assert created_study.study_name == loaded_study.study_name
@parametrize_storages
def test_load_study_study_name_none(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
if storage is None:
# `InMemoryStorage` can not be used with `load_study` function.
return
study_name = str(uuid.uuid4())
_ = create_study(study_name=study_name, storage=storage)
loaded_study = load_study(study_name=None, storage=storage)
assert loaded_study.study_name == study_name
study_name = str(uuid.uuid4())
_ = create_study(study_name=study_name, storage=storage)
# Ambiguous study.
with pytest.raises(ValueError):
load_study(study_name=None, storage=storage)
@parametrize_storages
def test_delete_study(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
# Test deleting a non-existing study.
with pytest.raises(KeyError):
delete_study(study_name="invalid-study-name", storage=storage)
# Test deleting an existing study.
study = create_study(storage=storage, load_if_exists=False)
delete_study(study_name=study.study_name, storage=storage)
# Test failed to delete the study which is already deleted.
with pytest.raises(KeyError):
delete_study(study_name=study.study_name, storage=storage)
def test_copy_study() -> None:
with StorageSupplier("sqlite") as from_storage, StorageSupplier("sqlite") as to_storage:
from_study = create_study(storage=from_storage)
from_study.set_user_attr("baz", "qux")
for _ in range(3):
trial = from_study.ask()
trial.suggest_float("x", 0, 1)
from_study.mark_comparison_ready(trial)
from_study.report_preference(from_study.trials[0], from_study.trials[1])
from_study.report_preference(from_study.trials[1], from_study.trials[2])
copy_study(
from_study_name=from_study.study_name,
from_storage=from_storage,
to_storage=to_storage,
)
to_study = load_study(study_name=from_study.study_name, storage=to_storage)
assert to_study.study_name == from_study.study_name
assert to_study.user_attrs == from_study.user_attrs
assert len(to_study.trials) == len(from_study.trials)
assert len(from_study.preferences) == len(to_study.preferences)
def test_copy_study_to_study_name() -> None:
with StorageSupplier("sqlite") as from_storage, StorageSupplier("sqlite") as to_storage:
from_study = create_study(study_name="foo", storage=from_storage)
_ = create_study(study_name="foo", storage=to_storage)
with pytest.raises(DuplicatedStudyError):
copy_study(
from_study_name=from_study.study_name,
from_storage=from_storage,
to_storage=to_storage,
)
copy_study(
from_study_name=from_study.study_name,
from_storage=from_storage,
to_storage=to_storage,
to_study_name="bar",
)
_ = load_study(study_name="bar", storage=to_storage)
@parametrize_storages
def test_add_trial(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
study = create_study(storage=storage)
assert len(study.trials) == 0
trial = create_trial(value=0)
study.add_trial(trial)
assert len(study.trials) == 1
assert study.trials[0].number == 0
def test_add_trial_invalid_values_length() -> None:
study = create_study()
trial = create_trial(values=[0, 0])
with pytest.raises(ValueError):
study.add_trial(trial)
@parametrize_storages
def test_add_trials(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
study = create_study(storage=storage)
assert len(study.trials) == 0
study.add_trials([])
assert len(study.trials) == 0
trials = [create_trial(value=i) for i in range(3)]
study.add_trials(trials)
assert len(study.trials) == 3
for i, trial in enumerate(study.trials):
assert trial.number == i
assert trial.value == i
other_study = create_study(storage=storage)
other_study.add_trials(study.trials)
assert len(other_study.trials) == 3
for i, trial in enumerate(other_study.trials):
assert trial.number == i
assert trial.value == i
@parametrize_storages
def test_get_trials(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
study = create_study(storage=storage)
for _ in range(5):
trial = study.ask()
trial.suggest_int("x", 1, 5)
study.mark_comparison_ready(trial)
with patch("copy.deepcopy", wraps=copy.deepcopy) as mock_object:
trials0 = study.get_trials(deepcopy=False)
assert mock_object.call_count == 0
assert len(trials0) == 5
trials1 = study.get_trials(deepcopy=True)
assert mock_object.call_count > 0
assert trials0 == trials1
# `study.trials` is equivalent to `study.get_trials(deepcopy=True)`.
old_count = mock_object.call_count
trials2 = study.trials
assert mock_object.call_count > old_count
assert trials0 == trials2
@parametrize_storages
def test_get_trials_state_option(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
study = create_study(storage=storage)
for _ in range(3):
trial = study.ask()
study.mark_comparison_ready(trial)
better, worse = study.trials[:2]
study.report_preference(better, worse)
trials = study.get_trials(states=None)
assert len(trials) == 3
trials = study.get_trials(states=(TrialState.RUNNING,))
assert len(trials) == 1
assert all(t.state == TrialState.RUNNING for t in trials)
trials = study.get_trials(states=(TrialState.COMPLETE,))
assert len(trials) == 2
assert all(t.state == TrialState.COMPLETE for t in trials)
trials = study.get_trials(states=())
assert len(trials) == 0
other_states = [
s for s in list(TrialState) if s != TrialState.COMPLETE and s != TrialState.RUNNING
]
for s in other_states:
trials = study.get_trials(states=(s,))
assert len(trials) == 0
def test_ask() -> None:
study = create_study()
trial = study.ask()
assert isinstance(trial, Trial)
def test_ask_fixed_search_space() -> None:
fixed_distributions = {
"x": distributions.FloatDistribution(0, 1),
"y": distributions.CategoricalDistribution(["bacon", "spam"]),
}
study = create_study()
trial = study.ask(fixed_distributions=fixed_distributions)
params = trial.params
assert len(trial.params) == 2
assert 0 <= params["x"] < 1
assert params["y"] in ["bacon", "spam"]
def test_report_preferences_from_another_process() -> None:
pool = multiprocessing.Pool()
with StorageSupplier("sqlite") as storage:
# Create a study and ask for a new trial.
study = create_study(storage=storage)
study.ask()
study.ask()
# Test normal behaviour.
better, worse = study.trials
pool.starmap(study.report_preference, [(better, worse)])
assert len(study.trials) == 2
assert study.trials[0].state == TrialState.COMPLETE
assert study.trials[1].state == TrialState.COMPLETE
assert len(study.preferences) == 1
@@ -0,0 +1,28 @@
from __future__ import annotations
from typing import Callable
import optuna
from optuna_dashboard.preferential._system_attrs import get_preferences
from optuna_dashboard.preferential._system_attrs import report_preferences
from ..storage_supplier import parametrize_storages
from ..storage_supplier import StorageSupplier
@parametrize_storages
def test_report_and_get_preferences(storage_supplier: Callable[[], StorageSupplier]) -> None:
with storage_supplier() as storage:
study = optuna.create_study(storage=storage)
study.ask()
study.ask()
assert len(get_preferences(study)) == 0
better, worse = study.trials[0], study.trials[1]
report_preferences(study, [(better, worse)])
assert len(get_preferences(study)) == 1
actual_better, actual_worse = get_preferences(study)[0]
assert actual_better.number == better.number
assert actual_worse.number == worse.number
+75
View File
@@ -0,0 +1,75 @@
from __future__ import annotations
import tempfile
from types import TracebackType
from typing import Any
from typing import IO
import optuna
from optuna.version import __version__ as optuna_ver
from packaging import version
import pytest
parametrize_storages = pytest.mark.parametrize(
"storage_supplier",
[
lambda: StorageSupplier("inmemory"),
lambda: StorageSupplier("sqlite"),
lambda: StorageSupplier("cached_sqlite"),
# TODO(c-bata): Support "JournalRedisStorage"
pytest.param(
lambda: StorageSupplier("journal"),
marks=pytest.mark.skipif(
version.parse(optuna_ver) < version.Version("3.1.0"),
reason="Artifact is not implemented yet in Optuna",
),
),
],
)
SQLITE3_TIMEOUT = 300
class StorageSupplier:
def __init__(self, storage_specifier: str, **kwargs: Any) -> None:
self.storage_specifier = storage_specifier
self.extra_args = kwargs
self.tempfile: IO[Any] | None = None
def __enter__(
self,
) -> (
optuna.storages.InMemoryStorage
| optuna.storages._CachedStorage
| optuna.storages.RDBStorage
| optuna.storages.JournalStorage
):
if self.storage_specifier == "inmemory":
if len(self.extra_args) > 0:
raise ValueError("InMemoryStorage does not accept any arguments!")
return optuna.storages.InMemoryStorage()
elif "sqlite" in self.storage_specifier:
self.tempfile = tempfile.NamedTemporaryFile(**self.extra_args)
url = "sqlite:///{}".format(self.tempfile.name)
rdb_storage = optuna.storages.RDBStorage(
url,
engine_kwargs={"connect_args": {"timeout": SQLITE3_TIMEOUT}},
**self.extra_args,
)
return (
optuna.storages._CachedStorage(rdb_storage)
if "cached" in self.storage_specifier
else rdb_storage
)
elif "journal" in self.storage_specifier:
self.tempfile = tempfile.NamedTemporaryFile(**self.extra_args)
file_storage = optuna.storages.JournalFileStorage(self.tempfile.name)
return optuna.storages.JournalStorage(file_storage)
else:
assert False, "Must not reach here"
def __exit__(
self, exc_type: type[BaseException], exc_val: BaseException, exc_tb: TracebackType
) -> None:
if self.tempfile:
self.tempfile.close()
+7 -4
View File
@@ -19,6 +19,9 @@ import { Clear } from "@mui/icons-material"
type Order = "asc" | "desc"
// eslint-disable-next-line @typescript-eslint/no-explicit-any
type Value = any
const defaultRowsPerPageOption = [10, 50, 100, { label: "All", value: -1 }]
interface DataGridColumn<T> {
@@ -33,7 +36,7 @@ interface DataGridColumn<T> {
interface RowFilter {
columnIdx: number
value: any
value: Value
}
function DataGrid<T>(props: {
@@ -45,7 +48,7 @@ function DataGrid<T>(props: {
initialRowsPerPage?: number
rowsPerPageOption?: Array<number | { value: number; label: string }>
defaultFilter?: (row: T) => boolean
}) {
}): React.ReactElement {
const { columns, rows, keyField, dense, collapseBody, defaultFilter } = props
let { initialRowsPerPage, rowsPerPageOption } = props
const [order, setOrder] = React.useState<Order>("asc")
@@ -81,7 +84,7 @@ function DataGrid<T>(props: {
const fieldAlreadyFiltered = (columnIdx: number): boolean =>
filters.some((f) => f.columnIdx === columnIdx)
const handleClickFilterCell = (columnIdx: number, value: any) => {
const handleClickFilterCell = (columnIdx: number, value: Value) => {
if (fieldAlreadyFiltered(columnIdx)) {
return
}
@@ -242,7 +245,7 @@ function DataGridRow<T>(props: {
row: T
keyField: keyof T
collapseBody?: (rowIndex: number) => React.ReactNode
handleClickFilterCell: (columnIdx: number, value: any) => void
handleClickFilterCell: (columnIdx: number, value: Value) => void
}) {
const {
columns,
+27 -11
View File
@@ -231,8 +231,23 @@ const plotHistory = (
return xAxis === "number"
? trial.number
: xAxis === "datetime_start"
? trial.datetime_start!
: trial.datetime_complete!
? trial.datetime_start ?? new Date()
: trial.datetime_complete ?? new Date()
}
const getValue = (trial: Trial, objectiveId: number): number | null => {
if (
objectiveId === null ||
trial.values === undefined ||
trial.values.length <= objectiveId
) {
return null
}
const value = trial.values[objectiveId]
if (value === "inf" || value === "-inf") {
return null
}
return value
}
const xForLinePlot: (number | Date)[] = []
@@ -240,34 +255,35 @@ const plotHistory = (
let currentBest: number | null = null
for (let i = 0; i < filteredTrials.length; i++) {
const t = filteredTrials[i]
const v = getValue(t, objectiveId) as number
if (currentBest === null) {
currentBest = t.values![objectiveId] as number
currentBest = v
xForLinePlot.push(getAxisX(t))
yForLinePlot.push(t.values![objectiveId] as number)
yForLinePlot.push(v)
} else if (
study.directions[objectiveId] === "maximize" &&
t.values![objectiveId] > currentBest
v > currentBest
) {
const p = filteredTrials[i - 1]
if (!xForLinePlot.includes(getAxisX(p))) {
xForLinePlot.push(getAxisX(p))
yForLinePlot.push(currentBest)
}
currentBest = t.values![objectiveId] as number
currentBest = v
xForLinePlot.push(getAxisX(t))
yForLinePlot.push(t.values![objectiveId] as number)
yForLinePlot.push(v)
} else if (
study.directions[objectiveId] === "minimize" &&
t.values![objectiveId] < currentBest
v < currentBest
) {
const p = filteredTrials[i - 1]
if (!xForLinePlot.includes(getAxisX(p))) {
xForLinePlot.push(getAxisX(p))
yForLinePlot.push(currentBest)
}
currentBest = t.values![objectiveId] as number
currentBest = v
xForLinePlot.push(getAxisX(t))
yForLinePlot.push(t.values![objectiveId] as number)
yForLinePlot.push(v)
}
}
xForLinePlot.push(getAxisX(filteredTrials[filteredTrials.length - 1]))
@@ -277,7 +293,7 @@ const plotHistory = (
{
x: filteredTrials.map(getAxisX),
y: filteredTrials.map(
(t: Trial): number => t.values![objectiveId] as number
(t: Trial): number => getValue(t, objectiveId) as number
),
name: "Objective Value",
mode: "markers",
+7
View File
@@ -7,9 +7,11 @@ export const loadStorage = (
setter: SetterOrUpdater<Study[]>
): void => {
sqlite3InitModule({
// eslint-disable-next-line @typescript-eslint/no-explicit-any
print: (...args: any): void => {
console.log(args)
},
// eslint-disable-next-line @typescript-eslint/no-explicit-any
printErr: (...args: any): void => {
console.log(args)
},
@@ -32,6 +34,7 @@ export const loadStorage = (
let supported = true
db.exec({
sql: "SELECT schema_version FROM version_info LIMIT 1",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
if (vals[0] != 12) {
supported = false
@@ -49,6 +52,7 @@ export const loadStorage = (
"SELECT s.study_id, s.study_name, sd.direction, sd.objective" +
" FROM studies AS s INNER JOIN study_directions AS sd" +
" ON s.study_id = sd.study_id ORDER BY sd.study_direction_id",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const study_id = vals[0]
const study_name = vals[1]
@@ -82,6 +86,7 @@ export const loadStorage = (
" FROM trials AS t LEFT JOIN trial_values AS tv ON tv.trial_id = t.trial_id" +
` WHERE t.study_id = ${s.study_id}` +
" ORDER BY t.number",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const state: TrialState =
vals[3] === "COMPLETE"
@@ -115,6 +120,7 @@ export const loadStorage = (
sql:
"SELECT param_name, param_value" +
` FROM trial_params WHERE trial_id = ${trial.trial_id}`,
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const param_name = vals[0]
params.push({
@@ -151,6 +157,7 @@ export const loadStorage = (
"SELECT value, value_type" +
` FROM trial_values WHERE trial_id = ${trial.trial_id}` +
" ORDER BY objective",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
values.push(
vals[1] === "INF_NEG"
-1
View File
@@ -10,7 +10,6 @@
],
"rules": {
"@typescript-eslint/naming-convention": "warn",
"@typescript-eslint/semi": "warn",
"curly": "warn",
"eqeqeq": "warn",
"no-throw-literal": "warn",