From ded02009049cba838e89641f5df04706d81dff47 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sat, 7 Jan 2023 21:13:16 +0900 Subject: [PATCH] Implement user API for Artifacts --- optuna_dashboard/_app.py | 9 +- optuna_dashboard/_serializer.py | 8 +- optuna_dashboard/artifact/__init__.py | 132 ++++++++++++++++++++++ optuna_dashboard/artifact/file_system.py | 27 +++++ pyproject.toml | 1 + python_tests/test_file_system_artifact.py | 22 ++++ 6 files changed, 196 insertions(+), 3 deletions(-) create mode 100644 optuna_dashboard/artifact/__init__.py create mode 100644 optuna_dashboard/artifact/file_system.py create mode 100644 python_tests/test_file_system_artifact.py diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 2330cf78..e7c010d6 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -36,6 +36,7 @@ from optuna.version import __version__ as optuna_ver from packaging import version from . import _note as note +from . import artifact from ._cached_extra_study_property import get_cached_extra_study_property from ._importance import get_param_importance_from_trials_cache from ._pareto_front import get_pareto_front_trials @@ -201,7 +202,11 @@ def get_trials(storage: BaseStorage, study_id: int, ttl_seconds: int = 10) -> li return trials -def create_app(storage: BaseStorage, debug: bool = False) -> Bottle: +def create_app( + storage: BaseStorage, + artifact_backend: Optional[artifact.ArtifactBackend] = None, + debug: bool = False, +) -> Bottle: app = Bottle() update_schema_compatibility_flags(storage) @@ -448,6 +453,8 @@ def create_app(storage: BaseStorage, debug: bool = False) -> Bottle: filename = gz_filename return static_file(filename, root=STATIC_DIR) + if artifact_backend is not None: + artifact.register_artifact_route(app, storage, artifact_backend) return app diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 6a8baf26..36567983 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -12,6 +12,7 @@ from optuna.study import StudySummary from optuna.trial import FrozenTrial from . import _note as note +from . import artifact from ._named_objectives import get_objective_names @@ -162,7 +163,7 @@ def serialize_frozen_trial( "distribution": serialize_distribution(distribution), } ) - trial_system_attrs = getattr(trial, "_system_attrs", {}) + trial_system_attrs: dict[str, Any] = getattr(trial, "_system_attrs", {}) fixed_params = trial_system_attrs.get("fixed_params", {}) serialized = { "trial_id": trial._trial_id, @@ -175,8 +176,11 @@ def serialize_frozen_trial( for param_name in fixed_params ], "user_attrs": serialize_attrs(trial.user_attrs), - "system_attrs": serialize_attrs(trial_system_attrs), + "system_attrs": serialize_attrs( + {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": artifact._list_artifacts(study_system_attrs, trial._trial_id), } serialized_intermediate_values: list[IntermediateValue] = [] diff --git a/optuna_dashboard/artifact/__init__.py b/optuna_dashboard/artifact/__init__.py new file mode 100644 index 00000000..06cf2721 --- /dev/null +++ b/optuna_dashboard/artifact/__init__.py @@ -0,0 +1,132 @@ +from __future__ import annotations + +import json +import mimetypes +import os.path +from typing import Any +from typing import BinaryIO +from typing import Optional +import uuid + +from bottle import Bottle +from bottle import response +import optuna +from optuna.storages import BaseStorage + + +try: + from typing import Protocol + from typing import TypedDict +except ImportError: + from typing_extensions import Protocol # type: ignore + from typing_extensions import TypedDict # type: ignore + + +ARTIFACTS_ATTR_PREFIX = "dashboard:artifacts:" + + +ArtifactMeta = TypedDict( + "ArtifactMeta", + { + "artifact_id": str, + "mimetype": str, + "encoding": str, + "filename": str, + }, +) + + +class ArtifactBackend(Protocol): + def open(self, artifact_id: str) -> BinaryIO: + ... + + def write(self, artifact_id: str, content_body: BinaryIO) -> None: + ... + + +def register_artifact_route( + app: Bottle, storage: BaseStorage, artifact_backend: ArtifactBackend +) -> None: + @app.get("/artifacts//") + def proxy_artifact(trial_id: int, artifact_id: str) -> bytes: + if artifact_backend is None: + response.status = 400 # Bad Request + return b"Cannot access to the artifacts." + artifact_dict = _get_artifact_meta(storage, trial_id, artifact_id) + response.set_header("Content-Type", artifact_dict["mimetype"]) + response.set_header("Content-Encodings", artifact_dict["encoding"]) + with artifact_backend.open(artifact_id) as f: + body = f.read() + return body + + +def upload_artifact( + backend: ArtifactBackend, + trial: optuna.Trial, + file_path: str, + *, + mimetype: Optional[str] = None, + encoding: Optional[str] = None, +) -> str: + """Upload an artifact (files), which is associated with the trial. + + Example: + .. code-block:: python + + import optuna + from optuna_dashboard.artifact import upload_artifact + from optuna_dashboard.artifact.file_system import FileSystemBackend + + artifact_backend = FileSystemBackend("./tmp/") + + def objective(trial: optuna.Trial) -> float: + ... = trial.suggest_float("x", -10, 10) + file_path = generate_example_png(...) + upload_artifact(artifact_backend, trial, file_path) + return ... + """ + filename = os.path.basename(file_path) + + guess_mimetype, guess_encoding = mimetypes.guess_type(filename) + mimetype = mimetype or guess_mimetype + encoding = encoding or guess_encoding + if mimetype is None or encoding is None: + raise ValueError("Failed to guess mimetype and encoding. Please explicitly specify them.") + + storage = trial.storage + trial_id = trial._trial_id + artifact_id = str(uuid.uuid4()) + artifact: ArtifactMeta = { + "artifact_id": artifact_id, + "mimetype": mimetype, + "encoding": encoding, + "filename": filename, + } + attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id + storage.set_study_system_attr(trial_id, attr_key, json.dumps(artifact)) + + with open(file_path, "rb") as f: + backend.write(artifact_id, f) + return artifact_id + + +def _artifact_prefix(trial_id: int) -> str: + return ARTIFACTS_ATTR_PREFIX + f"{trial_id}:" + + +def _get_artifact_meta(storage: BaseStorage, trial_id: int, artifact_id: str) -> ArtifactMeta: + artifact_key = ARTIFACTS_ATTR_PREFIX + artifact_id + storage.get_trial_system_attrs(trial_id) + + for key, value in storage.get_trial_system_attrs(trial_id).items(): + if key == artifact_key: + return json.loads(value) + raise ValueError("Artifact not found") + + +def _list_artifacts(study_system_attrs: dict[str, Any], trial_id: int) -> list[ArtifactMeta]: + return [ + json.loads(value) + for key, value in study_system_attrs.items() + if key.startswith(_artifact_prefix(trial_id)) + ] diff --git a/optuna_dashboard/artifact/file_system.py b/optuna_dashboard/artifact/file_system.py new file mode 100644 index 00000000..97d74ef8 --- /dev/null +++ b/optuna_dashboard/artifact/file_system.py @@ -0,0 +1,27 @@ +from __future__ import annotations + +import os +from typing import BinaryIO +from typing import TYPE_CHECKING + +from . import ArtifactBackend + + +class FileSystemBackend: + def __init__(self, base_path: str) -> None: + self._base_path = base_path + + def open(self, artifact_id: str) -> BinaryIO: + filepath = os.path.join(self._base_path, artifact_id) + return open(filepath, "rb") + + def write(self, artifact_id: str, content_body: BinaryIO) -> None: + filepath = os.path.join(self._base_path, artifact_id) + with open(filepath, "wb") as f: + f.write(content_body.read()) + + +if TYPE_CHECKING: + # A mypy-runtime assertion to ensure that LocalArtifactBackend + # implements all abstract methods in ArtifactBackendProtocol. + _: ArtifactBackend = FileSystemBackend("") diff --git a/pyproject.toml b/pyproject.toml index 96d29c75..51283cb0 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -30,6 +30,7 @@ dependencies = [ "optuna>=2.4.0", "packaging", "scikit-learn", + 'typing-extensions; python_version<"3.8"', ] dynamic = ["version"] diff --git a/python_tests/test_file_system_artifact.py b/python_tests/test_file_system_artifact.py new file mode 100644 index 00000000..572a88cd --- /dev/null +++ b/python_tests/test_file_system_artifact.py @@ -0,0 +1,22 @@ +import io +import tempfile +from unittest import TestCase + +from optuna_dashboard.artifact.file_system import FileSystemBackend + + +class FileSystemBackendTestCase(TestCase): + def setUp(self) -> None: + self.dir = tempfile.TemporaryDirectory() + + def tearDown(self) -> None: + self.dir.cleanup() + + def test_upload_download(self) -> None: + artifact_id = "dummy-uuid" + dummy_content = b"Hello World" + backend = FileSystemBackend(self.dir.name) + backend.write(artifact_id, io.BytesIO(dummy_content)) + with backend.open(artifact_id) as f: + actual = f.read() + self.assertEqual(actual, dummy_content)