From 81d21eadab1daa1d7631931f46d1a7eadc24f22d Mon Sep 17 00:00:00 2001 From: c-bata Date: Fri, 8 Sep 2023 17:20:01 +0900 Subject: [PATCH] Add support for Optuna's study artifact --- optuna_dashboard/_serializer.py | 2 + optuna_dashboard/artifact/_backend.py | 82 +++++++++++++++---- optuna_dashboard/ts/apiClient.ts | 2 + optuna_dashboard/ts/types/index.d.ts | 1 + python_tests/artifact/test_backend.py | 10 +-- .../artifact/test_optuna_compatibility.py | 4 +- 6 files changed, 79 insertions(+), 22 deletions(-) diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 19acbbd5..4f6f8dc3 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -16,6 +16,7 @@ from . import _note as note from ._form_widget import get_form_widgets_json from ._named_objectives import get_objective_names from ._preferential_history import _SYSTEM_ATTR_PREFIX_HISTORY +from .artifact._backend import list_study_artifacts from .artifact._backend import list_trial_artifacts from .preferential._study import _SYSTEM_ATTR_PREFERENTIAL_STUDY @@ -140,6 +141,7 @@ def serialize_study_detail( "user_attrs": serialize_attrs(summary.user_attrs), } system_attrs = getattr(summary, "system_attrs", {}) + serialized["artifacts"] = list_study_artifacts(system_attrs) if summary.datetime_start is not None: serialized["datetime_start"] = summary.datetime_start.isoformat() diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index b83a8a0c..63fec1c5 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -49,24 +49,49 @@ BaseRequest.MEMFILE_MAX = int( def get_artifact_path( - trial: optuna.Trial, + study_or_trial: optuna.Trial | optuna.Study, artifact_id: str, ) -> str: """Get the URL path for a given artifact ID.""" - study_id = trial.study._study_id - trial_id = trial._trial_id + if isinstance(study_or_trial, optuna.Study): + study_id = study_or_trial._study_id + return f"/artifacts/{study_id}/{artifact_id}" + + study_id = study_or_trial.study._study_id + trial_id = study_or_trial._trial_id return f"/artifacts/{study_id}/{trial_id}/{artifact_id}" def register_artifact_route( app: Bottle, storage: BaseStorage, artifact_store: Optional[ArtifactStore] ) -> None: - @app.get("/artifacts///") - def proxy_artifact(study_id: int, trial_id: int, artifact_id: str) -> HTTPResponse | bytes: + @app.get("/artifacts//") + def proxy_study_artifact(study_id: int, artifact_id: str) -> HTTPResponse | bytes: 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_study_artifact_meta(storage, study_id, artifact_id) + if artifact_dict is None: + response.status = 404 + return b"Not Found" + headers = {"Content-Type": artifact_dict["mimetype"]} + encoding = artifact_dict.get("encoding") + if encoding: + headers["Content-Encodings"] = encoding + + fp = artifact_store.open_reader(artifact_id) + return HTTPResponse(fp, headers=headers) + + @app.get("/artifacts///") + def proxy_trial_artifact( + study_id: int, + trial_id: int, + artifact_id: str, + ) -> HTTPResponse | bytes: + if artifact_store is None: + response.status = 400 # Bad Request + return b"Cannot access to the artifacts." + artifact_dict = get_trial_artifact_meta(storage, study_id, trial_id, artifact_id) if artifact_dict is None: response.status = 404 return b"Not Found" @@ -129,7 +154,7 @@ def register_artifact_route( # The artifact's metadata is stored in one of the following two locations: storage.set_study_system_attr( - study_id, _artifact_prefix(trial_id) + artifact_id, json.dumps(None) + study_id, _dashboard_trial_artifact_prefix(trial_id) + artifact_id, json.dumps(None) ) storage.set_trial_system_attr( trial_id, ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None) @@ -141,7 +166,7 @@ def register_artifact_route( def upload_artifact( backend: ArtifactBackend, - trial: optuna.Trial, + study_or_trial: optuna.Trial | optuna.Study, file_path: str, *, mimetype: Optional[str] = None, @@ -177,8 +202,6 @@ def upload_artifact( ) filename = os.path.basename(file_path) - storage = trial.storage - trial_id = trial._trial_id artifact_id = str(uuid.uuid4()) guess_mimetype, guess_encoding = mimetypes.guess_type(filename) artifact: ArtifactMeta = { @@ -188,23 +211,42 @@ def upload_artifact( "filename": filename, } attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id - storage.set_trial_system_attr(trial_id, attr_key, json.dumps(artifact)) + + if isinstance(study_or_trial, optuna.Study): + storage = study_or_trial._storage + study_id = study_or_trial._study_id + storage.set_study_system_attr(study_id, attr_key, json.dumps(artifact)) + else: + storage = study_or_trial.storage + trial_id = study_or_trial._trial_id + storage.set_trial_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: +def _dashboard_trial_artifact_prefix(trial_id: int) -> str: return DASHBOARD_ARTIFACTS_ATTR_PREFIX + f"{trial_id}:" -def get_artifact_meta( +def get_study_artifact_meta( + storage: BaseStorage, study_id: int, artifact_id: str +) -> Optional[ArtifactMeta]: + study_system_attrs = storage.get_study_system_attrs(study_id) + attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id + artifact_meta = study_system_attrs.get(attr_key) + if artifact_meta is not None: + return json.loads(artifact_meta) + return None + + +def get_trial_artifact_meta( storage: BaseStorage, study_id: int, trial_id: int, artifact_id: str ) -> Optional[ArtifactMeta]: # Search study_system_attrs due to backward compatibility. study_system_attrs = storage.get_study_system_attrs(study_id) - attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id + attr_key = _dashboard_trial_artifact_prefix(trial_id=trial_id) + artifact_id artifact_meta = study_system_attrs.get(attr_key) if artifact_meta is not None: return json.loads(artifact_meta) @@ -223,6 +265,7 @@ def get_artifact_meta( def delete_all_artifacts(backend: ArtifactStore, storage: BaseStorage, study_id: int) -> None: artifact_metas = [] study_system_attrs = storage.get_study_system_attrs(study_id) + artifact_metas.extend(list_study_artifacts(study_system_attrs)) for trial in storage.get_all_trials(study_id): trial_artifacts = list_trial_artifacts(study_system_attrs, trial) artifact_metas.extend(trial_artifacts) @@ -231,6 +274,15 @@ def delete_all_artifacts(backend: ArtifactStore, storage: BaseStorage, study_id: backend.remove(meta["artifact_id"]) +def list_study_artifacts(study_system_attrs: dict[str, Any]) -> list[ArtifactMeta]: + artifact_metas = [ + json.loads(value) + for key, value in study_system_attrs.items() + if key.startswith(ARTIFACTS_ATTR_PREFIX) + ] + return [a for a in artifact_metas if a is not None] + + def list_trial_artifacts( study_system_attrs: dict[str, Any], trial: FrozenTrial ) -> list[ArtifactMeta]: @@ -238,7 +290,7 @@ def list_trial_artifacts( dashboard_artifact_metas = [ json.loads(value) for key, value in study_system_attrs.items() - if key.startswith(_artifact_prefix(trial._trial_id)) + if key.startswith(_dashboard_trial_artifact_prefix(trial._trial_id)) ] # Collect ArtifactMeta from trial_system_attrs. Note that artifacts uploaded via diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index e62e0e42..afe1379f 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -94,6 +94,7 @@ interface StudyDetailResponse { form_widgets?: FormWidgets preference_history?: PreferenceHistoryResponce[] plotly_graph_objects: PlotlyGraphObject[] + artifacts: Artifact[] } export const getStudyDetailAPI = ( @@ -133,6 +134,7 @@ export const getStudyDetailAPI = ( convertPreferenceHistory ), plotly_graph_objects: res.data.plotly_graph_objects, + artifacts: res.data.artifacts, } }) } diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index b7b35797..df3798bb 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -205,6 +205,7 @@ type StudyDetail = { form_widgets?: FormWidgets preference_history?: PreferenceHistory[] plotly_graph_objects: PlotlyGraphObject[] + artifacts: Artifact[] } type StudyDetails = { diff --git a/python_tests/artifact/test_backend.py b/python_tests/artifact/test_backend.py index e2d0f152..74698a09 100644 --- a/python_tests/artifact/test_backend.py +++ b/python_tests/artifact/test_backend.py @@ -8,11 +8,11 @@ import pytest def test_get_artifact_path() -> None: study = MagicMock(_study_id=0) trial = MagicMock(_trial_id=0, study=study) - assert _backend.get_artifact_path(trial=trial, artifact_id="id0") == "/artifacts/0/0/id0" + assert _backend.get_artifact_path(trial, "id0") == "/artifacts/0/0/id0" def test_artifact_prefix() -> None: - actual = _backend._artifact_prefix(trial_id=0) + actual = _backend._dashboard_trial_artifact_prefix(trial_id=0) assert actual == "dashboard:artifacts:0:" @@ -47,13 +47,13 @@ def init_storage_with_artifact_meta() -> BaseStorage: def test_get_artifact_meta(init_storage_with_artifact_meta: MagicMock) -> None: storage = init_storage_with_artifact_meta - actual = _backend.get_artifact_meta(storage, study_id=0, trial_id=0, artifact_id="id0") + actual = _backend.get_trial_artifact_meta(storage, study_id=0, trial_id=0, artifact_id="id0") assert actual == {"artifact_id": "id0", "filename": "foo.txt"} - actual = _backend.get_artifact_meta(storage, study_id=0, trial_id=1, artifact_id="id3") + actual = _backend.get_trial_artifact_meta(storage, study_id=0, trial_id=1, artifact_id="id3") assert actual == {"artifact_id": "id3", "filename": "qux.txt"} - actual = _backend.get_artifact_meta(storage, study_id=0, trial_id=0, artifact_id="id4") + actual = _backend.get_trial_artifact_meta(storage, study_id=0, trial_id=0, artifact_id="id4") assert actual is None diff --git a/python_tests/artifact/test_optuna_compatibility.py b/python_tests/artifact/test_optuna_compatibility.py index d81e1397..82d36450 100644 --- a/python_tests/artifact/test_optuna_compatibility.py +++ b/python_tests/artifact/test_optuna_compatibility.py @@ -6,7 +6,7 @@ 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 get_trial_artifact_meta from optuna_dashboard.artifact._backend import list_trial_artifacts from packaging import version import pytest @@ -44,7 +44,7 @@ def test_list_optuna_trial_artifacts() -> None: with artifact_store.open_reader(artifact_id) as reader: assert reader.read() == dummy_content - artifact_meta = get_artifact_meta( + artifact_meta = get_trial_artifact_meta( storage=storage, study_id=study._study_id, trial_id=trial._trial_id,