From a6cfb00c65f3915acf677ff4acd4d7658db876bb Mon Sep 17 00:00:00 2001 From: c-bata Date: Fri, 9 Feb 2024 16:27:38 +0900 Subject: [PATCH] Revert "Merge pull request #719 from nabenabe0928/enhance/speedup-get-trials" This reverts commit cc26f0e72fa63f100244d6b5d1209ccba70bee56, reversing changes made to 83b548f9327497f2e1fb379d43bdefe896cb4993. --- optuna_dashboard/_app.py | 27 +++---- optuna_dashboard/_serializer.py | 2 - optuna_dashboard/ts/action.ts | 18 +---- optuna_dashboard/ts/apiClient.ts | 6 +- optuna_dashboard/ts/components/AppDrawer.tsx | 7 +- .../ts/components/StudyDetail.tsx | 10 +-- optuna_dashboard/ts/state.ts | 5 -- optuna_dashboard/ts/types/index.d.ts | 1 - python_tests/test_api.py | 74 +++++++++++-------- python_tests/test_serializers.py | 4 +- setup.cfg | 3 - 11 files changed, 61 insertions(+), 96 deletions(-) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index b57edc13..ac0e1cf0 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -193,20 +193,14 @@ def create_app( @app.get("/api/studies/") @json_api_view def get_study_detail(study_id: int) -> dict[str, Any]: - # Use the following default values if not specified in request.params. - query_params = dict(after=0, limit=2000) - for query_key in query_params: - try: - query_params[query_key] = int(request.params[query_key]) - assert query_params[query_key] >= 0 - except AssertionError: - response.status = 400 # Bad parameter - return {"reason": f"`{query_key}` should be larger than or equal to 0."} - except KeyError: - # Use the default parameter defined in query_params. - pass - - after, limit = query_params["after"], query_params["limit"] + 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 @@ -237,19 +231,16 @@ def create_app( plotly_graph_objects = get_plotly_graph_objects(system_attrs) skipped_trial_ids = get_skipped_trial_ids(system_attrs) skipped_trial_numbers = [t.number for t in trials if t._trial_id in skipped_trial_ids] - limit = len(trials) if limit == 0 else limit - fetched_trials_partially = after + limit < len(trials) return serialize_study_detail( summary, best_trials, - trials[after : after + limit], + trials[after:], intersection, union, union_user_attrs, has_intermediate_values, plotly_graph_objects, skipped_trial_numbers, - fetched_trials_partially, ) @app.get("/api/studies//param_importances") diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index d279ef5a..e3c4da65 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -143,13 +143,11 @@ def serialize_study_detail( has_intermediate_values: bool, plotly_graph_objects: dict[str, str], skipped_trial_numbers: list[int], - fetched_trials_partially: bool, ) -> dict[str, Any]: serialized: dict[str, Any] = { "name": summary.study_name, "directions": [d.name.lower() for d in summary.directions], "user_attrs": serialize_attrs(summary.user_attrs), - "fetched_trials_partially": fetched_trials_partially, } system_attrs = getattr(summary, "system_attrs", {}) serialized["artifacts"] = list_study_artifacts(system_attrs) diff --git a/optuna_dashboard/ts/action.ts b/optuna_dashboard/ts/action.ts index 7c27a84c..0c448c2e 100644 --- a/optuna_dashboard/ts/action.ts +++ b/optuna_dashboard/ts/action.ts @@ -27,7 +27,6 @@ import { studySummariesState, paramImportanceState, isFileUploading, - fetchedTrialsPartiallyState, artifactIsAvailable, plotlypyIsAvailableState, reloadIntervalState, @@ -49,9 +48,6 @@ export const actionCreator = () => { const setUploading = useSetRecoilState(isFileUploading) const setTrialsUpdating = useSetRecoilState(trialsUpdatingState) const setArtifactIsAvailable = useSetRecoilState(artifactIsAvailable) - const setFetchedTrialsPartially = useSetRecoilState( - fetchedTrialsPartiallyState - ) const setPlotlypyIsAvailable = useSetRecoilState( plotlypyIsAvailableState ) @@ -247,12 +243,8 @@ export const actionCreator = () => { }) } - const updateStudyDetail = ( - studyId: number, - forceFetchAllTrials: boolean = false - ) => { + const updateStudyDetail = (studyId: number) => { let nLocalFixedTrials = 0 - const nMaximumTrialsAtOnce = forceFetchAllTrials ? 0 : 2000 if (studyId in studyDetails) { const currentTrials = studyDetails[studyId].trials const firstUpdatable = currentTrials.findIndex((trial) => @@ -261,20 +253,14 @@ export const actionCreator = () => { nLocalFixedTrials = firstUpdatable === -1 ? currentTrials.length : firstUpdatable } - getStudyDetailAPI(studyId, nLocalFixedTrials, nMaximumTrialsAtOnce) + getStudyDetailAPI(studyId, nLocalFixedTrials) .then((study) => { - if (studyId in studyDetails && study.trials.length === 0) { - // Update trials only if necessary. - // NOTE: The first condition is for study with no trials. - return - } const currentFixedTrials = studyId in studyDetails ? studyDetails[studyId].trials.slice(0, nLocalFixedTrials) : [] study.trials = currentFixedTrials.concat(study.trials) setStudyDetailState(studyId, study) - setFetchedTrialsPartially(study.fetched_trials_partially) }) .catch((err) => { const reason = err.response?.data.reason diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index 87f7409a..14566a07 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -104,19 +104,16 @@ interface StudyDetailResponse { artifacts: Artifact[] feedback_component_type: FeedbackComponentType skipped_trial_numbers?: number[] - fetched_trials_partially: boolean } export const getStudyDetailAPI = ( studyId: number, - nLocalTrials: number, - nMaximumTrialsAtOnce: number + nLocalTrials: number ): Promise => { return axiosInstance .get(`/api/studies/${studyId}`, { params: { after: nLocalTrials, - limit: nMaximumTrialsAtOnce, }, }) .then((res) => { @@ -150,7 +147,6 @@ export const getStudyDetailAPI = ( plotly_graph_objects: res.data.plotly_graph_objects, artifacts: res.data.artifacts, skipped_trial_numbers: res.data.skipped_trial_numbers ?? [], - fetched_trials_partially: res.data.fetched_trials_partially, } }) } diff --git a/optuna_dashboard/ts/components/AppDrawer.tsx b/optuna_dashboard/ts/components/AppDrawer.tsx index db09bfcf..1d4d9c78 100644 --- a/optuna_dashboard/ts/components/AppDrawer.tsx +++ b/optuna_dashboard/ts/components/AppDrawer.tsx @@ -324,12 +324,7 @@ export const AppDrawer: FC<{ { - const newReloadInterval = reloadInterval === -1 ? 10 : -1 - action.saveReloadInterval(newReloadInterval) - if (newReloadInterval === -1) { - const forceFetchAllTrials = true - action.updateStudyDetail(studyId, forceFetchAllTrials) - } + action.saveReloadInterval(reloadInterval === -1 ? 10 : -1) }} > diff --git a/optuna_dashboard/ts/components/StudyDetail.tsx b/optuna_dashboard/ts/components/StudyDetail.tsx index 3477b03a..065b90f5 100644 --- a/optuna_dashboard/ts/components/StudyDetail.tsx +++ b/optuna_dashboard/ts/components/StudyDetail.tsx @@ -16,7 +16,6 @@ import HomeIcon from "@mui/icons-material/Home" import { StudyNote } from "./Note" import { actionCreator } from "../action" import { - fetchedTrialsPartiallyState, reloadIntervalState, useStudyDetailValue, useStudyIsPreferential, @@ -57,9 +56,6 @@ export const StudyDetail: FC<{ const reloadInterval = useRecoilValue(reloadIntervalState) const studyName = useStudyName(studyId) const isPreferential = useStudyIsPreferential(studyId) - const fetchedTrialsPartially = useRecoilValue( - fetchedTrialsPartiallyState - ) const title = studyName !== null ? `${studyName} (id=${studyId})` : `Study #${studyId}` @@ -76,13 +72,9 @@ export const StudyDetail: FC<{ const nTrials = studyDetail ? studyDetail.trials.length : 0 let interval = reloadInterval * 1000 - // If trials are left in cache, we collect them quickly. // For Human-in-the-loop Optimization, the interval is set to 2 seconds // when the number of trials is small, and the page is "trialList" or top page of preferential. - if (fetchedTrialsPartially) { - // Too short time is frustrating because the page freezes until the rendering is done. - interval = 3000 - } else if ( + if ( (!isPreferential && page === "trialList") || (isPreferential && page === "top") ) { diff --git a/optuna_dashboard/ts/state.ts b/optuna_dashboard/ts/state.ts index 8ae05cf2..8ca7586a 100644 --- a/optuna_dashboard/ts/state.ts +++ b/optuna_dashboard/ts/state.ts @@ -38,11 +38,6 @@ export const drawerOpenState = atom({ default: false, }) -export const fetchedTrialsPartiallyState = atom({ - key: "fetchedTrialsPartially", - default: false, -}) - export const isFileUploading = atom({ key: "isFileUploading", default: false, diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index c018a077..30d6708f 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -220,7 +220,6 @@ type StudyDetail = { plotly_graph_objects: PlotlyGraphObject[] artifacts: Artifact[] skipped_trial_numbers: number[] - fetched_trials_partially: boolean } type StudyDetails = { diff --git a/python_tests/test_api.py b/python_tests/test_api.py index 96035255..d644635f 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -2,7 +2,6 @@ from __future__ import annotations import json import sys -from typing import Any from unittest import TestCase import optuna @@ -46,53 +45,70 @@ class APITestCase(TestCase): study_summaries = json.loads(body)["study_summaries"] self.assertEqual(len(study_summaries), 2) - def run_get_study_details( - self, - queries: dict[str, str] | None = None, - expected_status: int = 200, - ) -> list[dict[str, Any]]: + def test_get_study_details_without_after_param(self) -> None: study = optuna.create_study() study_id = study._study_id - study.optimize(objective, n_trials=10) + study.optimize(objective, n_trials=2) app = create_app(study._storage) status, _, body = send_request( app, f"/api/studies/{study_id}", "GET", - queries=queries, content_type="application/json", ) - self.assertEqual(status, expected_status) - if expected_status == 400: - return [] - else: - return json.loads(body)["trials"] - - def test_get_study_details_without_after_param(self) -> None: - all_trials = self.run_get_study_details() - self.assertEqual(len(all_trials), 10) + self.assertEqual(status, 200) + all_trials = json.loads(body)["trials"] + self.assertEqual(len(all_trials), 2) def test_get_study_details_with_after_param_partial(self) -> None: - all_trials = self.run_get_study_details({"after": "5"}) - self.assertEqual(len(all_trials), 5) + study = optuna.create_study() + study_id = study._study_id + study.optimize(objective, n_trials=2) + app = create_app(study._storage) - def test_get_study_details_with_params(self) -> None: - for after in [0, 5, 9, 10]: - for limit in [1, 2, 5, 10]: - trials = self.run_get_study_details({"after": str(after), "limit": str(limit)}) - ans = list(range(after, min(10, after + limit))) - self.assertEqual([t["number"] for t in trials], ans) + status, _, body = send_request( + app, + f"/api/studies/{study_id}", + "GET", + queries={"after": "1"}, + content_type="application/json", + ) + self.assertEqual(status, 200) + all_trials = json.loads(body)["trials"] + self.assertEqual(len(all_trials), 1) def test_get_study_details_with_after_param_full(self) -> None: - all_trials = self.run_get_study_details({"after": "10"}) + study = optuna.create_study() + study_id = study._study_id + study.optimize(objective, n_trials=2) + app = create_app(study._storage) + + status, _, body = send_request( + app, + f"/api/studies/{study_id}", + "GET", + queries={"after": "2"}, + content_type="application/json", + ) + self.assertEqual(status, 200) + all_trials = json.loads(body)["trials"] self.assertEqual(len(all_trials), 0) def test_get_study_details_with_after_param_illegal(self) -> None: - self.run_get_study_details({"after": "-1"}, expected_status=400) + study = optuna.create_study() + study_id = study._study_id + study.optimize(objective, n_trials=2) + app = create_app(study._storage) - def test_get_study_details_with_limit_param_illegal(self) -> None: - self.run_get_study_details({"limit": "-1"}, expected_status=400) + status, _, body = send_request( + app, + f"/api/studies/{study_id}", + "GET", + queries={"after": "-1"}, + content_type="application/json", + ) + self.assertEqual(status, 400) @pytest.mark.skipif(sys.version_info < (3, 8), reason="BoTorch dropped Python3.7 support") @pytest.mark.skipif( diff --git a/python_tests/test_serializers.py b/python_tests/test_serializers.py index c37587c8..a991d5d3 100644 --- a/python_tests/test_serializers.py +++ b/python_tests/test_serializers.py @@ -65,7 +65,7 @@ def test_get_study_detail_is_preferential() -> None: study_summary = study_summaries[0] study_detail = serialize_study_detail( - study_summary, [], study.trials, [], [], [], False, {}, [], False + study_summary, [], study.trials, [], [], [], False, {}, [] ) assert study_detail["is_preferential"] @@ -78,7 +78,7 @@ def test_get_study_detail_is_not_preferential() -> None: study_summary = study_summaries[0] study_detail = serialize_study_detail( - study_summary, [], study.trials, [], [], [], False, {}, [], False + study_summary, [], study.trials, [], [], [], False, {}, [] ) assert not study_detail["is_preferential"] diff --git a/setup.cfg b/setup.cfg index 1eabbdce..93d9ad56 100644 --- a/setup.cfg +++ b/setup.cfg @@ -1,7 +1,4 @@ [flake8] -ignore = - E203 - W503 max-line-length = 99 statistics = True exclude = venv,build