diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index b7702c35..990706da 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -35,7 +35,7 @@ from ._storage import get_study_summaries from ._storage import get_study_summary from ._storage import get_trials from ._storage_url import get_storage -from .artifact._backend import delete_all_artifacts +from .artifact._backend import delete_all_study_artifacts from .artifact._backend import register_artifact_route from .artifact._backend_to_store import ArtifactBackendToStore from .artifact._backend_to_store import is_artifact_store @@ -161,8 +161,7 @@ def create_app( @json_api_view def delete_study(study_id: int) -> dict[str, Any]: if artifact_store is not None: - system_attrs = storage.get_study_system_attrs(study_id) - delete_all_artifacts(artifact_store, system_attrs) + delete_all_study_artifacts(artifact_store, storage, study_id) try: storage.delete_study(study_id) diff --git a/optuna_dashboard/artifact/_backend.py b/optuna_dashboard/artifact/_backend.py index 63ae8fb8..10297a89 100644 --- a/optuna_dashboard/artifact/_backend.py +++ b/optuna_dashboard/artifact/_backend.py @@ -66,7 +66,7 @@ def register_artifact_route( 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" @@ -81,6 +81,7 @@ def register_artifact_route( @app.post("/api/artifacts//") @json_api_view def upload_artifact_api(study_id: int, trial_id: int) -> dict[str, Any]: + # 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."} @@ -190,27 +191,38 @@ 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_study_artifacts( + backend: ArtifactBackend, storage: BaseStorage, study_id: int +) -> None: + for meta in list_study_artifacts(storage, study_id): backend.remove(meta["artifact_id"]) +def list_study_artifacts(storage: BaseStorage, study_id: int) -> list[ArtifactMeta]: + artifact_metas = [] + study_system_attrs = storage.get_study_system_attrs(study_id) + for trial in storage.get_all_trials(study_id): + artifact_metas.extend(list_trial_artifacts(study_system_attrs, trial)) + return artifact_metas + + def list_trial_artifacts( study_system_attrs: dict[str, Any], trial: FrozenTrial ) -> list[ArtifactMeta]: diff --git a/python_tests/artifact/test_optuna_compatibility.py b/python_tests/artifact/test_optuna_compatibility.py index 1c1114bf..7f6cd264 100644 --- a/python_tests/artifact/test_optuna_compatibility.py +++ b/python_tests/artifact/test_optuna_compatibility.py @@ -1,19 +1,22 @@ from __future__ import annotations +import os import tempfile import optuna from optuna.version import __version__ as optuna_ver +from optuna_dashboard.artifact._backend import get_artifact_meta +from optuna_dashboard.artifact._backend import list_study_artifacts 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"), + version.parse(optuna_ver) < version.Version("3.3.0.dev"), reason="Artifact is not implemented yet in Optuna", ) -def test_list_optuna_artifacts() -> None: +def test_list_optuna_trial_artifacts() -> None: from optuna.artifacts import FileSystemArtifactStore from optuna.artifacts import upload_artifact @@ -40,3 +43,39 @@ def test_list_optuna_artifacts() -> None: 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.dev"), + reason="Artifact is not implemented yet in Optuna", +) +def test_list_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) + artifact_meta_list = list_study_artifacts(storage, study_id=study._study_id) + assert len(artifact_meta_list) == 10