From 8ecec450a1c81e23fcd2488392ac72f4dc8f4d0c Mon Sep 17 00:00:00 2001 From: c-bata Date: Fri, 10 Feb 2023 10:30:48 +0900 Subject: [PATCH] Introduce --storage-class to support journal storages --- optuna_dashboard/_app.py | 14 +---- optuna_dashboard/_cli.py | 12 +++- optuna_dashboard/_storage_url.py | 96 ++++++++++++++++++++++++++++++++ 3 files changed, 106 insertions(+), 16 deletions(-) create mode 100644 optuna_dashboard/_storage_url.py diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 9a732298..97f58aa9 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -37,6 +37,7 @@ from ._importance import get_param_importance_from_trials_cache from ._pareto_front import get_pareto_front_trials from ._serializer import serialize_study_detail from ._serializer import serialize_study_summary +from ._storage_url import get_storage from .artifact._backend import delete_all_artifacts from .artifact._backend import register_artifact_route @@ -503,19 +504,6 @@ def _frozen_study_to_study_summary(frozen_study: "FrozenStudy") -> StudySummary: ) -def get_storage(storage: Union[str, BaseStorage]) -> BaseStorage: - if isinstance(storage, str): - if storage.startswith("redis"): - raise ValueError( - "RedisStorage is unsupported from Optuna v3.1 or Optuna Dashboard v0.8.0" - ) - elif version.parse(optuna_ver) >= version.Version("v3.0.0"): - return RDBStorage(storage, skip_compatibility_check=True, skip_table_creation=True) - else: - return RDBStorage(storage, skip_compatibility_check=True) - return storage - - def run_server( storage: Union[str, BaseStorage], host: str = "localhost", diff --git a/optuna_dashboard/_cli.py b/optuna_dashboard/_cli.py index 9e5749f4..16d290cc 100644 --- a/optuna_dashboard/_cli.py +++ b/optuna_dashboard/_cli.py @@ -15,8 +15,8 @@ from optuna.storages import RDBStorage from . import __version__ from ._app import create_app -from ._app import get_storage from ._sql_profiler import register_profiler_view +from ._storage_url import get_storage from .artifact.file_system import FileSystemBackend @@ -84,7 +84,13 @@ def auto_select_server( def main() -> None: parser = argparse.ArgumentParser(description="Real-time dashboard for Optuna.") - parser.add_argument("storage", help="DB URL (e.g. sqlite:///example.db)", type=str) + parser.add_argument("storage", help="Storage URL (e.g. sqlite:///example.db)", type=str) + parser.add_argument( + "--storage-class", + help="Storage class hint (e.g. JournalFileStorage)", + type=str, + default=None, + ) parser.add_argument( "--port", help="port number (default: %(default)s)", type=int, default=8080 ) @@ -105,7 +111,7 @@ def main() -> None: args = parser.parse_args() storage: BaseStorage - storage = get_storage(args.storage) + storage = get_storage(args.storage, storage_class=args.storage_class) artifact_backend = None if args.artifact_dir is not None: diff --git a/optuna_dashboard/_storage_url.py b/optuna_dashboard/_storage_url.py new file mode 100644 index 00000000..f67623c7 --- /dev/null +++ b/optuna_dashboard/_storage_url.py @@ -0,0 +1,96 @@ +from __future__ import annotations + +import os.path +import re +from typing import TYPE_CHECKING + +from optuna.storages import BaseStorage +from optuna.storages import RDBStorage +from optuna.version import __version__ as optuna_ver +from packaging import version + + +if TYPE_CHECKING: + from typing import Optional + from typing import Union + + from optuna.storages import JournalStorage + + +# https://github.com/zzzeek/sqlalchemy/blob/c6554ac52/lib/sqlalchemy/engine/url.py#L234-L292 +rfc1738_pattern = re.compile( + r""" + (?P[\w\+]+):// + (?: + (?P[^:/]*) + (?::(?P.*))? + @)? + (?: + (?: + \[(?P[^/]+)\] | + (?P[^/:]+) + )? + (?::(?P[^/]*))? + )? + (?:/(?P.*))? + """, + re.X, +) + + +def get_storage( + storage: Union[str, BaseStorage], storage_class: Optional[str] = None +) -> BaseStorage: + if isinstance(storage, BaseStorage): + return storage + + if storage_class: + if storage_class == "RDBStorage": + return get_rdb_storage(storage) + if storage_class == "JournalRedisStorage": + return get_journal_redis_storage(storage) + if storage_class == "JournalFileStorage": + return get_journal_file_storage(storage) + raise ValueError("Unexpected storage_class") + + return guess_storage_from_url(storage) + + +def guess_storage_from_url(storage_url: str) -> BaseStorage: + if storage_url.startswith("redis"): + return get_journal_redis_storage(storage_url) + + if os.path.isfile(storage_url): + return get_journal_file_storage(storage_url) + + if rfc1738_pattern.match(storage_url) is not None: + return get_rdb_storage(storage_url) + + raise ValueError("Failed to guess storage class from storage_url") + + +def get_rdb_storage(storage_url: str) -> RDBStorage: + if version.parse(optuna_ver) >= version.Version("v3.0.0"): + return RDBStorage(storage_url, skip_compatibility_check=True, skip_table_creation=True) + else: + return RDBStorage(storage_url, skip_compatibility_check=True) + + +def get_journal_file_storage(file_path: str) -> JournalStorage: + if version.parse(optuna_ver) < version.Version("v3.1.0"): + raise ValueError("JournalRedisStorage is available from Optuna v3.1.0") + + from optuna.storages import JournalFileStorage + from optuna.storages import JournalStorage + + return JournalStorage(JournalFileStorage(file_path=file_path)) + + +def get_journal_redis_storage(redis_url: str) -> JournalStorage: + if version.parse(optuna_ver) < version.Version("v3.1.0"): + raise ValueError("JournalRedisStorage is available from Optuna v3.1.0") + + from optuna.storages import JournalRedisStorage + from optuna.storages import JournalStorage + + return JournalStorage(JournalRedisStorage(redis_url))