From a7429318f3231ae6713cd7d9954a5c8e9a1b8548 Mon Sep 17 00:00:00 2001 From: nabenabe0928 Date: Fri, 1 Dec 2023 10:22:11 +0100 Subject: [PATCH] Limit the maximum number of trials to be used for one update More specifically, this PR makes the following changes: 1. Add query param for limit, 2. Update studyDetails only if there is no study with the specified study_id or the fetched trials has a positive length, 3. Shorten the waiting interval when there is a leftover in the server side, and 4. Adapt the flake8 setup to the Optuna repo. --- optuna_dashboard/_app.py | 25 ++++++++++++------- optuna_dashboard/_serializer.py | 2 ++ optuna_dashboard/ts/action.ts | 14 ++++++++--- optuna_dashboard/ts/apiClient.ts | 6 ++++- .../ts/components/StudyDetail.tsx | 8 +++++- optuna_dashboard/ts/state.ts | 5 ++++ optuna_dashboard/ts/types/index.d.ts | 1 + python_tests/test_serializers.py | 4 +-- setup.cfg | 3 +++ 9 files changed, 52 insertions(+), 16 deletions(-) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index ead27188..1dd05854 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -191,14 +191,19 @@ 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 + query_params = dict(after=0, limit=1000) + 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 @@ -229,16 +234,18 @@ 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] + 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 7030abec..3bc659c3 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 10fe43bc..3947a0ee 100644 --- a/optuna_dashboard/ts/action.ts +++ b/optuna_dashboard/ts/action.ts @@ -27,6 +27,7 @@ import { studySummariesState, paramImportanceState, isFileUploading, + isTrialLeftInCache, artifactIsAvailable, reloadIntervalState, trialsUpdatingState, @@ -46,6 +47,7 @@ export const actionCreator = () => { const setUploading = useSetRecoilState(isFileUploading) const setTrialsUpdating = useSetRecoilState(trialsUpdatingState) const setArtifactIsAvailable = useSetRecoilState(artifactIsAvailable) + const setIsTrialLeftInCache = useSetRecoilState(isTrialLeftInCache) const setStudyDetailState = (studyId: number, study: StudyDetail) => { setStudyDetails((prevVal) => { @@ -233,6 +235,7 @@ export const actionCreator = () => { const updateStudyDetail = (studyId: number) => { let nLocalFixedTrials = 0 + let nMaximumTrialsAtOnce = 1000 if (studyId in studyDetails) { const currentTrials = studyDetails[studyId].trials const firstUpdatable = currentTrials.findIndex((trial) => @@ -240,15 +243,20 @@ export const actionCreator = () => { ) nLocalFixedTrials = firstUpdatable === -1 ? currentTrials.length : firstUpdatable + nMaximumTrialsAtOnce = 2000 } - getStudyDetailAPI(studyId, nLocalFixedTrials) + getStudyDetailAPI(studyId, nLocalFixedTrials, nMaximumTrialsAtOnce) .then((study) => { const currentFixedTrials = studyId in studyDetails ? studyDetails[studyId].trials.slice(0, nLocalFixedTrials) : [] - study.trials = currentFixedTrials.concat(study.trials) - setStudyDetailState(studyId, study) + if (study.trials.length !== 0 || !(studyId in studyDetails)) { + // Update trials only if necessary. The second condition is for study with no trials. + study.trials = currentFixedTrials.concat(study.trials) + setStudyDetailState(studyId, study) + setIsTrialLeftInCache(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 f42d20de..ec80056d 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -102,16 +102,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) => { @@ -145,6 +148,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/StudyDetail.tsx b/optuna_dashboard/ts/components/StudyDetail.tsx index 2f56ec8b..1f2823dd 100644 --- a/optuna_dashboard/ts/components/StudyDetail.tsx +++ b/optuna_dashboard/ts/components/StudyDetail.tsx @@ -17,6 +17,7 @@ import DownloadIcon from "@mui/icons-material/Download" import { StudyNote } from "./Note" import { actionCreator } from "../action" import { + isTrialLeftInCache, reloadIntervalState, useStudyDetailValue, useStudyIsPreferential, @@ -57,6 +58,7 @@ export const StudyDetail: FC<{ const reloadInterval = useRecoilValue(reloadIntervalState) const studyName = useStudyName(studyId) const isPreferential = useStudyIsPreferential(studyId) + const isTrialLeft = useRecoilValue(isTrialLeftInCache) const title = studyName !== null ? `${studyName} (id=${studyId})` : `Study #${studyId}` @@ -73,9 +75,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 (isTrialLeft) { + // 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 b1f60654..7825311b 100644 --- a/optuna_dashboard/ts/state.ts +++ b/optuna_dashboard/ts/state.ts @@ -33,6 +33,11 @@ export const drawerOpenState = atom({ default: false, }) +export const isTrialLeftInCache = atom({ + key: "isTrialLeftInCache", + 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_serializers.py b/python_tests/test_serializers.py index d1bdf59b..1f52d926 100644 --- a/python_tests/test_serializers.py +++ b/python_tests/test_serializers.py @@ -61,7 +61,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"] @@ -74,7 +74,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