diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index ac0e1cf0..b57edc13 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -193,14 +193,20 @@ def create_app( @app.get("/api/studies/") @json_api_view def get_study_detail(study_id: int) -> dict[str, Any]: - 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 + # 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"] summary = get_study_summary(storage, study_id) if summary is None: response.status = 404 # Not found @@ -231,16 +237,19 @@ 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:], + trials[after : after + limit], 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 b280d411..07291819 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -141,11 +141,13 @@ 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 0c448c2e..7c27a84c 100644 --- a/optuna_dashboard/ts/action.ts +++ b/optuna_dashboard/ts/action.ts @@ -27,6 +27,7 @@ import { studySummariesState, paramImportanceState, isFileUploading, + fetchedTrialsPartiallyState, artifactIsAvailable, plotlypyIsAvailableState, reloadIntervalState, @@ -48,6 +49,9 @@ export const actionCreator = () => { const setUploading = useSetRecoilState(isFileUploading) const setTrialsUpdating = useSetRecoilState(trialsUpdatingState) const setArtifactIsAvailable = useSetRecoilState(artifactIsAvailable) + const setFetchedTrialsPartially = useSetRecoilState( + fetchedTrialsPartiallyState + ) const setPlotlypyIsAvailable = useSetRecoilState( plotlypyIsAvailableState ) @@ -243,8 +247,12 @@ export const actionCreator = () => { }) } - const updateStudyDetail = (studyId: number) => { + const updateStudyDetail = ( + studyId: number, + forceFetchAllTrials: boolean = false + ) => { let nLocalFixedTrials = 0 + const nMaximumTrialsAtOnce = forceFetchAllTrials ? 0 : 2000 if (studyId in studyDetails) { const currentTrials = studyDetails[studyId].trials const firstUpdatable = currentTrials.findIndex((trial) => @@ -253,14 +261,20 @@ export const actionCreator = () => { nLocalFixedTrials = firstUpdatable === -1 ? currentTrials.length : firstUpdatable } - getStudyDetailAPI(studyId, nLocalFixedTrials) + getStudyDetailAPI(studyId, nLocalFixedTrials, nMaximumTrialsAtOnce) .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 14566a07..87f7409a 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -104,16 +104,19 @@ interface StudyDetailResponse { artifacts: Artifact[] feedback_component_type: FeedbackComponentType skipped_trial_numbers?: number[] + fetched_trials_partially: boolean } export const getStudyDetailAPI = ( studyId: number, - nLocalTrials: number + nLocalTrials: number, + nMaximumTrialsAtOnce: number ): Promise => { return axiosInstance .get(`/api/studies/${studyId}`, { params: { after: nLocalTrials, + limit: nMaximumTrialsAtOnce, }, }) .then((res) => { @@ -147,6 +150,7 @@ 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 31eeec69..0c17f338 100644 --- a/optuna_dashboard/ts/components/AppDrawer.tsx +++ b/optuna_dashboard/ts/components/AppDrawer.tsx @@ -310,7 +310,12 @@ export const AppDrawer: FC<{ { - action.saveReloadInterval(reloadInterval === -1 ? 10 : -1) + const newReloadInterval = reloadInterval === -1 ? 10 : -1 + action.saveReloadInterval(newReloadInterval) + if (newReloadInterval === -1) { + const forceFetchAllTrials = true + action.updateStudyDetail(studyId, forceFetchAllTrials) + } }} > diff --git a/optuna_dashboard/ts/components/StudyDetail.tsx b/optuna_dashboard/ts/components/StudyDetail.tsx index 065b90f5..3477b03a 100644 --- a/optuna_dashboard/ts/components/StudyDetail.tsx +++ b/optuna_dashboard/ts/components/StudyDetail.tsx @@ -16,6 +16,7 @@ import HomeIcon from "@mui/icons-material/Home" import { StudyNote } from "./Note" import { actionCreator } from "../action" import { + fetchedTrialsPartiallyState, reloadIntervalState, useStudyDetailValue, useStudyIsPreferential, @@ -56,6 +57,9 @@ 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}` @@ -72,9 +76,13 @@ 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 ( + if (fetchedTrialsPartially) { + // Too short time is frustrating because the page freezes until the rendering is done. + interval = 3000 + } else if ( (!isPreferential && page === "trialList") || (isPreferential && page === "top") ) { diff --git a/optuna_dashboard/ts/state.ts b/optuna_dashboard/ts/state.ts index 20f75c2d..eff4f26d 100644 --- a/optuna_dashboard/ts/state.ts +++ b/optuna_dashboard/ts/state.ts @@ -34,6 +34,11 @@ 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 67182ff8..b1ed1bb7 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -220,6 +220,7 @@ 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 d644635f..96035255 100644 --- a/python_tests/test_api.py +++ b/python_tests/test_api.py @@ -2,6 +2,7 @@ from __future__ import annotations import json import sys +from typing import Any from unittest import TestCase import optuna @@ -45,70 +46,53 @@ class APITestCase(TestCase): study_summaries = json.loads(body)["study_summaries"] self.assertEqual(len(study_summaries), 2) - def test_get_study_details_without_after_param(self) -> None: + def run_get_study_details( + self, + queries: dict[str, str] | None = None, + expected_status: int = 200, + ) -> list[dict[str, Any]]: study = optuna.create_study() study_id = study._study_id - study.optimize(objective, n_trials=2) + study.optimize(objective, n_trials=10) 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, 200) - all_trials = json.loads(body)["trials"] - self.assertEqual(len(all_trials), 2) + 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) def test_get_study_details_with_after_param_partial(self) -> None: - study = optuna.create_study() - study_id = study._study_id - study.optimize(objective, n_trials=2) - app = create_app(study._storage) + all_trials = self.run_get_study_details({"after": "5"}) + self.assertEqual(len(all_trials), 5) - 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_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) def test_get_study_details_with_after_param_full(self) -> None: - 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"] + all_trials = self.run_get_study_details({"after": "10"}) self.assertEqual(len(all_trials), 0) def test_get_study_details_with_after_param_illegal(self) -> None: - study = optuna.create_study() - study_id = study._study_id - study.optimize(objective, n_trials=2) - app = create_app(study._storage) + self.run_get_study_details({"after": "-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) + def test_get_study_details_with_limit_param_illegal(self) -> None: + self.run_get_study_details({"limit": "-1"}, expected_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 a991d5d3..c37587c8 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, {}, [] + study_summary, [], study.trials, [], [], [], False, {}, [], 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, {}, [] + study_summary, [], study.trials, [], [], [], False, {}, [], False ) assert not study_detail["is_preferential"] diff --git a/setup.cfg b/setup.cfg index 93d9ad56..1eabbdce 100644 --- a/setup.cfg +++ b/setup.cfg @@ -1,4 +1,7 @@ [flake8] +ignore = + E203 + W503 max-line-length = 99 statistics = True exclude = venv,build