From c79135b52411db6f46310e89a6a62c111ba34fa5 Mon Sep 17 00:00:00 2001 From: c-bata Date: Fri, 18 Feb 2022 14:59:09 +0900 Subject: [PATCH] Deprecate create_app() function --- optuna_dashboard/__init__.py | 4 +- optuna_dashboard/_app.py | 313 ++++++++++++++++++ optuna_dashboard/{cli.py => _cli.py} | 2 +- .../{search_space.py => _search_space.py} | 0 .../{serializer.py => _serializer.py} | 0 optuna_dashboard/app.py | 312 +---------------- python_tests/test_api.py | 2 +- python_tests/test_search_space.py | 2 +- python_tests/test_serializers.py | 2 +- setup.cfg | 2 +- 10 files changed, 327 insertions(+), 312 deletions(-) create mode 100644 optuna_dashboard/_app.py rename optuna_dashboard/{cli.py => _cli.py} (98%) rename optuna_dashboard/{search_space.py => _search_space.py} (100%) rename optuna_dashboard/{serializer.py => _serializer.py} (100%) diff --git a/optuna_dashboard/__init__.py b/optuna_dashboard/__init__.py index ebd9110f..5ddc00a6 100644 --- a/optuna_dashboard/__init__.py +++ b/optuna_dashboard/__init__.py @@ -1,5 +1,5 @@ -from .app import run_server # noqa -from .app import wsgi # noqa +from ._app import run_server # noqa +from ._app import wsgi # noqa __version__ = "0.5.0" diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py new file mode 100644 index 00000000..cafdff30 --- /dev/null +++ b/optuna_dashboard/_app.py @@ -0,0 +1,313 @@ +from datetime import datetime +from datetime import timedelta +import functools +import json +import logging +import os +import threading +import traceback +import typing +from typing import Any +from typing import Callable +from typing import cast +from typing import Dict +from typing import List +from typing import NoReturn +from typing import Optional +from typing import TypeVar +from typing import Union + +from bottle import BaseResponse +from bottle import Bottle +from bottle import redirect +from bottle import request +from bottle import response +from bottle import run +from bottle import static_file +import optuna +from optuna.exceptions import DuplicatedStudyError +from optuna.storages import BaseStorage +from optuna.storages import RDBStorage +from optuna.storages import RedisStorage +from optuna.study import Study +from optuna.study import StudyDirection +from optuna.study import StudySummary +from optuna.trial import FrozenTrial +from optuna.trial import TrialState + +from ._search_space import get_search_space +from ._serializer import serialize_study_detail +from ._serializer import serialize_study_summary + + +if typing.TYPE_CHECKING: + from _typeshed.wsgi import WSGIApplication + +BottleViewReturn = Union[str, bytes, Dict[str, Any], BaseResponse] +BottleView = TypeVar("BottleView", bound=Callable[..., BottleViewReturn]) + +logger = logging.getLogger(__name__) + +BASE_DIR = os.path.dirname(os.path.abspath(__file__)) +STATIC_DIR = os.path.join(BASE_DIR, "public") +INDEX_HTML = """ + + + Optuna Dashboard + + + + + + +
+

Now loading...

+
+ + +""" +# In-memory trials cache +trials_cache_lock = threading.Lock() +trials_cache: Dict[int, List[FrozenTrial]] = {} +trials_last_fetched_at: Dict[int, datetime] = {} + + +def handle_json_api_exception(view: BottleView) -> BottleView: + @functools.wraps(view) + def decorated(*args: List[Any], **kwargs: Dict[str, Any]) -> BottleViewReturn: + try: + response_body = view(*args, **kwargs) + return response_body + except Exception as e: + response.status = 500 + response.content_type = "application/json" + stacktrace = "\n".join(traceback.format_tb(e.__traceback__)) + logger.error(f"Exception: {e}\n{stacktrace}") + return json.dumps({"reason": "internal server error"}) + + return cast(BottleView, decorated) + + +def get_study_summary(storage: BaseStorage, study_id: int) -> Optional[StudySummary]: + summaries = storage.get_all_study_summaries() + for summary in summaries: + if summary._study_id != study_id: + continue + return summary + return None + + +def get_trials( + storage: BaseStorage, study_id: int, ttl_seconds: int = 10 +) -> List[FrozenTrial]: + with trials_cache_lock: + trials = trials_cache.get(study_id, None) + last_fetched_at = trials_last_fetched_at.get(study_id, None) + if ( + trials is not None + and last_fetched_at is not None + and datetime.now() - last_fetched_at < timedelta(seconds=ttl_seconds) + ): + return trials + trials = storage.get_all_trials(study_id) + with trials_cache_lock: + trials_last_fetched_at[study_id] = datetime.now() + trials_cache[study_id] = trials + return trials + + +def get_distribution_name(param_name: str, study: Study) -> str: + for trial in study.trials: + if param_name in trial.distributions: + return trial.distributions[param_name].__class__.__name__ + assert False, "Must not reach here." + + +def create_app(storage: BaseStorage) -> Bottle: + app = Bottle() + + @app.hook("before_request") + def remove_trailing_slashes_hook() -> None: + request.environ["PATH_INFO"] = request.environ["PATH_INFO"].rstrip("/") + + @app.get("/") + def index() -> BottleViewReturn: + return redirect("/dashboard", 302) # Status Found + + # Accept any following paths for client-side routing + @app.get("/dashboard<:re:(/.*)?>") + def dashboard() -> BottleViewReturn: + response.content_type = "text/html" + return INDEX_HTML + + @app.get("/api/studies") + @handle_json_api_exception + def list_study_summaries() -> BottleViewReturn: + response.content_type = "application/json" + summaries = [ + serialize_study_summary(summary) + for summary in storage.get_all_study_summaries() + ] + return { + "study_summaries": summaries, + } + + @app.post("/api/studies") + @handle_json_api_exception + def create_study() -> BottleViewReturn: + response.content_type = "application/json" + + study_name = request.json.get("study_name", None) + directions = request.json.get("directions", []) + if ( + study_name is None + or len(directions) == 0 + or not all([d in ("minimize", "maximize") for d in directions]) + ): + response.status = 400 # Bad request + return {"reason": "You need to set study_name and direction"} + + try: + study_id = storage.create_new_study(study_name) + except DuplicatedStudyError: + response.status = 400 # Bad request + return {"reason": f"'{study_name}' is already exists"} + + storage.set_study_directions( + study_id, + [ + StudyDirection.MAXIMIZE + if d.lower() == "maximize" + else StudyDirection.MINIMIZE + for d in directions + ], + ) + + summary = get_study_summary(storage, study_id) + if summary is None: + response.status = 500 # Internal server error + return {"reason": "Failed to create study"} + response.status = 201 # Created + return {"study_summary": serialize_study_summary(summary)} + + @app.delete("/api/studies/") + @handle_json_api_exception + def delete_study(study_id: int) -> BottleViewReturn: + response.content_type = "application/json" + + try: + storage.delete_study(study_id) + except KeyError: + response.status = 404 # Not found + return {"reason": f"study_id={study_id} is not found"} + response.status = 204 # No content + return "" + + @app.get("/api/studies/") + @handle_json_api_exception + def get_study_detail(study_id: int) -> BottleViewReturn: + response.content_type = "application/json" + try: + after = int(request.params["after"]) + assert after >= 0 + except AssertionError: + response.status = 400 # Bad parameter + return {"reason": "`after` should be larger or equal 0."} + except KeyError: + after = 0 + summary = get_study_summary(storage, study_id) + if summary is None: + response.status = 404 # Not found + return {"reason": f"study_id={study_id} is not found"} + trials = get_trials(storage, study_id)[after:] + intersection, union = get_search_space(study_id, trials) + return serialize_study_detail(summary, trials, intersection, union) + + @app.get("/api/studies//param_importances") + @handle_json_api_exception + def get_param_importances(study_id: int) -> BottleViewReturn: + # TODO(chenghuzi): add support for selecting params via query parameters. + response.content_type = "application/json" + objective_id = int(request.params.get("objective_id", 0)) + try: + study_name = storage.get_study_name_from_id(study_id) + study = Study(study_name=study_name, storage=storage) + except KeyError: + response.status = 404 # Not found + return {"reason": f"study_id={study_id} is not found"} + + n_directions = len(study.directions) + if objective_id >= n_directions: + response.status = 400 # Bad request + return { + "reason": f"study_id={study_id} has only {n_directions} direction(s)." + } + + completed_trials = [ + trial for trial in study.trials if trial.state == TrialState.COMPLETE + ] + evaluator = None + params = None + + if len(completed_trials) > 0: + importances = optuna.importance.get_param_importances( + study, + evaluator=evaluator, + params=params, + target=lambda t: t.values[objective_id], + ) + else: + importances = {} + target_name = "Objective Value" + + return { + "target_name": target_name, + "param_importances": [ + { + "name": name, + "importance": importance, + "distribution": get_distribution_name(name, study), + } + for name, importance in importances.items() + ], + } + + @app.get("/static/") + def send_static(filename: str) -> BottleViewReturn: + return static_file(filename, root=STATIC_DIR) + + return app + + +def get_storage(storage: Union[str, BaseStorage]) -> BaseStorage: + if isinstance(storage, str): + if storage.startswith("redis"): + return RedisStorage(storage) + else: + return RDBStorage(storage) + return storage + + +def run_server( # type: ignore + storage: Union[str, BaseStorage], host: str = "localhost", port: int = 8080 +) -> NoReturn: + """Start running optuna-dashboard and blocks until the server terminates. + 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 `wsgi()` function. + """ + app = create_app(get_storage(storage)) + run(app, host=host, port=port) + + +def wsgi(storage: Union[str, BaseStorage]) -> "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)) diff --git a/optuna_dashboard/cli.py b/optuna_dashboard/_cli.py similarity index 98% rename from optuna_dashboard/cli.py rename to optuna_dashboard/_cli.py index 87f43dfe..7812de0d 100644 --- a/optuna_dashboard/cli.py +++ b/optuna_dashboard/_cli.py @@ -9,7 +9,7 @@ from optuna.storages import RDBStorage from optuna.storages import RedisStorage from . import __version__ -from .app import create_app +from ._app import create_app AUTO_RELOAD = os.environ.get("OPTUNA_DASHBOARD_AUTO_RELOAD") == "1" diff --git a/optuna_dashboard/search_space.py b/optuna_dashboard/_search_space.py similarity index 100% rename from optuna_dashboard/search_space.py rename to optuna_dashboard/_search_space.py diff --git a/optuna_dashboard/serializer.py b/optuna_dashboard/_serializer.py similarity index 100% rename from optuna_dashboard/serializer.py rename to optuna_dashboard/_serializer.py diff --git a/optuna_dashboard/app.py b/optuna_dashboard/app.py index 70094be5..18ecfe38 100644 --- a/optuna_dashboard/app.py +++ b/optuna_dashboard/app.py @@ -1,312 +1,14 @@ -from datetime import datetime -from datetime import timedelta -import functools -import json -import logging -import os -import threading -import traceback -import typing -from typing import Any -from typing import Callable -from typing import cast -from typing import Dict -from typing import List -from typing import NoReturn -from typing import Optional -from typing import TypeVar -from typing import Union +import warnings -from bottle import BaseResponse from bottle import Bottle -from bottle import redirect -from bottle import request -from bottle import response -from bottle import run -from bottle import static_file -import optuna -from optuna.exceptions import DuplicatedStudyError from optuna.storages import BaseStorage -from optuna.storages import RDBStorage -from optuna.storages import RedisStorage -from optuna.study import Study -from optuna.study import StudyDirection -from optuna.study import StudySummary -from optuna.trial import FrozenTrial -from optuna.trial import TrialState -from . import serializer -from .search_space import get_search_space - - -if typing.TYPE_CHECKING: - from _typeshed.wsgi import WSGIApplication - -BottleViewReturn = Union[str, bytes, Dict[str, Any], BaseResponse] -BottleView = TypeVar("BottleView", bound=Callable[..., BottleViewReturn]) - -logger = logging.getLogger(__name__) - -BASE_DIR = os.path.dirname(os.path.abspath(__file__)) -STATIC_DIR = os.path.join(BASE_DIR, "public") -INDEX_HTML = """ - - - Optuna Dashboard - - - - - - -
-

Now loading...

-
- - -""" -# In-memory trials cache -trials_cache_lock = threading.Lock() -trials_cache: Dict[int, List[FrozenTrial]] = {} -trials_last_fetched_at: Dict[int, datetime] = {} - - -def handle_json_api_exception(view: BottleView) -> BottleView: - @functools.wraps(view) - def decorated(*args: List[Any], **kwargs: Dict[str, Any]) -> BottleViewReturn: - try: - response_body = view(*args, **kwargs) - return response_body - except Exception as e: - response.status = 500 - response.content_type = "application/json" - stacktrace = "\n".join(traceback.format_tb(e.__traceback__)) - logger.error(f"Exception: {e}\n{stacktrace}") - return json.dumps({"reason": "internal server error"}) - - return cast(BottleView, decorated) - - -def get_study_summary(storage: BaseStorage, study_id: int) -> Optional[StudySummary]: - summaries = storage.get_all_study_summaries() - for summary in summaries: - if summary._study_id != study_id: - continue - return summary - return None - - -def get_trials( - storage: BaseStorage, study_id: int, ttl_seconds: int = 10 -) -> List[FrozenTrial]: - with trials_cache_lock: - trials = trials_cache.get(study_id, None) - last_fetched_at = trials_last_fetched_at.get(study_id, None) - if ( - trials is not None - and last_fetched_at is not None - and datetime.now() - last_fetched_at < timedelta(seconds=ttl_seconds) - ): - return trials - trials = storage.get_all_trials(study_id) - with trials_cache_lock: - trials_last_fetched_at[study_id] = datetime.now() - trials_cache[study_id] = trials - return trials - - -def get_distribution_name(param_name: str, study: Study) -> str: - for trial in study.trials: - if param_name in trial.distributions: - return trial.distributions[param_name].__class__.__name__ - assert False, "Must not reach here." +from ._app import create_app as _create_app def create_app(storage: BaseStorage) -> Bottle: - app = Bottle() - - @app.hook("before_request") - def remove_trailing_slashes_hook() -> None: - request.environ["PATH_INFO"] = request.environ["PATH_INFO"].rstrip("/") - - @app.get("/") - def index() -> BottleViewReturn: - return redirect("/dashboard", 302) # Status Found - - # Accept any following paths for client-side routing - @app.get("/dashboard<:re:(/.*)?>") - def dashboard() -> BottleViewReturn: - response.content_type = "text/html" - return INDEX_HTML - - @app.get("/api/studies") - @handle_json_api_exception - def list_study_summaries() -> BottleViewReturn: - response.content_type = "application/json" - summaries = [ - serializer.serialize_study_summary(summary) - for summary in storage.get_all_study_summaries() - ] - return { - "study_summaries": summaries, - } - - @app.post("/api/studies") - @handle_json_api_exception - def create_study() -> BottleViewReturn: - response.content_type = "application/json" - - study_name = request.json.get("study_name", None) - directions = request.json.get("directions", []) - if ( - study_name is None - or len(directions) == 0 - or not all([d in ("minimize", "maximize") for d in directions]) - ): - response.status = 400 # Bad request - return {"reason": "You need to set study_name and direction"} - - try: - study_id = storage.create_new_study(study_name) - except DuplicatedStudyError: - response.status = 400 # Bad request - return {"reason": f"'{study_name}' is already exists"} - - storage.set_study_directions( - study_id, - [ - StudyDirection.MAXIMIZE - if d.lower() == "maximize" - else StudyDirection.MINIMIZE - for d in directions - ], - ) - - summary = get_study_summary(storage, study_id) - if summary is None: - response.status = 500 # Internal server error - return {"reason": "Failed to create study"} - response.status = 201 # Created - return {"study_summary": serializer.serialize_study_summary(summary)} - - @app.delete("/api/studies/") - @handle_json_api_exception - def delete_study(study_id: int) -> BottleViewReturn: - response.content_type = "application/json" - - try: - storage.delete_study(study_id) - except KeyError: - response.status = 404 # Not found - return {"reason": f"study_id={study_id} is not found"} - response.status = 204 # No content - return "" - - @app.get("/api/studies/") - @handle_json_api_exception - def get_study_detail(study_id: int) -> BottleViewReturn: - response.content_type = "application/json" - try: - after = int(request.params["after"]) - assert after >= 0 - except AssertionError: - response.status = 400 # Bad parameter - return {"reason": "`after` should be larger or equal 0."} - except KeyError: - after = 0 - summary = get_study_summary(storage, study_id) - if summary is None: - response.status = 404 # Not found - return {"reason": f"study_id={study_id} is not found"} - trials = get_trials(storage, study_id)[after:] - intersection, union = get_search_space(study_id, trials) - return serializer.serialize_study_detail(summary, trials, intersection, union) - - @app.get("/api/studies//param_importances") - @handle_json_api_exception - def get_param_importances(study_id: int) -> BottleViewReturn: - # TODO(chenghuzi): add support for selecting params via query parameters. - response.content_type = "application/json" - objective_id = int(request.params.get("objective_id", 0)) - try: - study_name = storage.get_study_name_from_id(study_id) - study = Study(study_name=study_name, storage=storage) - except KeyError: - response.status = 404 # Not found - return {"reason": f"study_id={study_id} is not found"} - - n_directions = len(study.directions) - if objective_id >= n_directions: - response.status = 400 # Bad request - return { - "reason": f"study_id={study_id} has only {n_directions} direction(s)." - } - - completed_trials = [ - trial for trial in study.trials if trial.state == TrialState.COMPLETE - ] - evaluator = None - params = None - - if len(completed_trials) > 0: - importances = optuna.importance.get_param_importances( - study, - evaluator=evaluator, - params=params, - target=lambda t: t.values[objective_id], - ) - else: - importances = {} - target_name = "Objective Value" - - return { - "target_name": target_name, - "param_importances": [ - { - "name": name, - "importance": importance, - "distribution": get_distribution_name(name, study), - } - for name, importance in importances.items() - ], - } - - @app.get("/static/") - def send_static(filename: str) -> BottleViewReturn: - return static_file(filename, root=STATIC_DIR) - - return app - - -def get_storage(storage: Union[str, BaseStorage]) -> BaseStorage: - if isinstance(storage, str): - if storage.startswith("redis"): - return RedisStorage(storage) - else: - return RDBStorage(storage) - return storage - - -def run_server( # type: ignore - storage: Union[str, BaseStorage], host: str = "localhost", port: int = 8080 -) -> NoReturn: - """Start running optuna-dashboard and blocks until the server terminates. - 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 `wsgi()` function. - """ - app = create_app(get_storage(storage)) - run(app, host=host, port=port) - - -def wsgi(storage: Union[str, BaseStorage]) -> "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)) + warnings.warn( + "This function will be removed in the future. Please use optuna_dashboard.run_server() instead.", + DeprecationWarning, + ) + return _create_app(storage) diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 0ee3c0ea..b3c82c6d 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -3,7 +3,7 @@ from unittest import TestCase import optuna -from optuna_dashboard.app import create_app +from optuna_dashboard._app import create_app from .wsgi_client import send_request diff --git a/python_tests/test_search_space.py b/python_tests/test_search_space.py index 597af5d9..df7fe9e2 100644 --- a/python_tests/test_search_space.py +++ b/python_tests/test_search_space.py @@ -7,7 +7,7 @@ from optuna.distributions import UniformDistribution from optuna.exceptions import ExperimentalWarning from optuna.trial import TrialState -from optuna_dashboard.search_space import _SearchSpace +from optuna_dashboard._search_space import _SearchSpace class SearchSpaceTestCase(TestCase): diff --git a/python_tests/test_serializers.py b/python_tests/test_serializers.py index c7e70f17..11de8080 100644 --- a/python_tests/test_serializers.py +++ b/python_tests/test_serializers.py @@ -1,6 +1,6 @@ from unittest import TestCase -from optuna_dashboard.serializer import serialize_attrs +from optuna_dashboard._serializer import serialize_attrs class SerializeAttrsTestCase(TestCase): diff --git a/setup.cfg b/setup.cfg index cd9f94e9..28087f69 100644 --- a/setup.cfg +++ b/setup.cfg @@ -35,7 +35,7 @@ install_requires = [options.entry_points] console_scripts = - optuna-dashboard = optuna_dashboard.cli:main + optuna-dashboard = optuna_dashboard._cli:main [options.package_data] optuna_dashboard = public/*