From 4fe8b36eaea207df51db5febe9eeaa279604b48a Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 1 Jan 2023 11:46:10 +0900 Subject: [PATCH 01/21] Refactor hyperparameter importance --- optuna_dashboard/_app.py | 11 ++-- optuna_dashboard/_importance.py | 42 +++++--------- optuna_dashboard/ts/action.ts | 30 ++++++++++ optuna_dashboard/ts/apiClient.ts | 19 ++---- .../GraphHyperparameterImportances.tsx | 58 ++++++++----------- optuna_dashboard/ts/state.ts | 5 ++ optuna_dashboard/ts/types/index.d.ts | 5 +- 7 files changed, 87 insertions(+), 83 deletions(-) diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 8183f48d..d10bc9c1 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -336,20 +336,19 @@ def create_app(storage: BaseStorage, debug: bool = False) -> Bottle: @app.get("/api/studies//param_importances") @json_api_view def get_param_importances(study_id: int) -> BottleViewReturn: - # TODO(chenghuzi): add support for selecting params via query parameters. - objective_id = int(request.params.get("objective_id", 0)) try: n_directions = len(storage.get_study_directions(study_id)) except KeyError: response.status = 404 # Study is not found return {"reason": f"study_id={study_id} is not found"} - if objective_id >= n_directions: - response.status = 400 # Bad request - return {"reason": f"study_id={study_id} has only {n_directions} direction(s)."} trials = get_trials(storage, study_id) try: - return get_param_importance_from_trials_cache(storage, study_id, objective_id, trials) + importances = [ + get_param_importance_from_trials_cache(storage, study_id, objective_id, trials) + for objective_id in range(n_directions) + ] + return {"param_importances": importances} except ValueError as e: response.status = 400 # Bad request return {"reason": str(e)} diff --git a/optuna_dashboard/_importance.py b/optuna_dashboard/_importance.py index f6e0e4fa..ec6507e1 100644 --- a/optuna_dashboard/_importance.py +++ b/optuna_dashboard/_importance.py @@ -1,6 +1,7 @@ from __future__ import annotations import threading +from typing import List from typing import TYPE_CHECKING import warnings @@ -25,26 +26,18 @@ except Exception as e: if TYPE_CHECKING: from typing import TypedDict - ImportanceItemType = TypedDict( - "ImportanceItemType", + ImportanceType = TypedDict( + "ImportanceType", { "name": str, "importance": float, "distribution": str, }, ) - ImportanceType = TypedDict( - "ImportanceType", - { - "target_name": str, - "param_importances": list[ImportanceItemType], - }, - ) -target_name = "Objective Value" param_importance_cache_lock = threading.Lock() # { "{study_id}:{objective_id}" : (n_completed_trials, importance) } -param_importance_cache: dict[str, tuple[int, ImportanceType]] = {} +param_importance_cache: dict[str, tuple[int, list[ImportanceType]]] = {} class StudyWrapper(Study): @@ -62,16 +55,16 @@ class StudyWrapper(Study): def get_param_importance_from_trials_cache( storage: BaseStorage, study_id: int, objective_id: int, trials: list[FrozenTrial] -) -> ImportanceType: +) -> list[ImportanceType]: completed_trials = [t for t in trials if t.state == TrialState.COMPLETE] n_completed_trials = len(completed_trials) if n_completed_trials == 0: - return {"target_name": target_name, "param_importances": []} + return [] cache_key = f"{study_id}:{objective_id}" with param_importance_cache_lock: cache_n_trial, cache_importance = param_importance_cache.get( - cache_key, (0, {"target_name": target_name, "param_importances": []}) + cache_key, (0, {"param_importances": []}) ) if n_completed_trials == cache_n_trial: return cache_importance @@ -95,18 +88,15 @@ def get_param_importance_from_trials_cache( def convert_to_importance_type( importance: dict[str, float], trials: list[FrozenTrial] -) -> ImportanceType: - return { - "target_name": target_name, - "param_importances": [ - { - "name": name, - "importance": importance, - "distribution": get_distribution_name(name, trials), - } - for name, importance in importance.items() - ], - } +) -> list[ImportanceType]: + return [ + { + "name": name, + "importance": importance, + "distribution": get_distribution_name(name, trials), + } + for name, importance in importance.items() + ] def get_distribution_name(param_name: str, trials: list[FrozenTrial]) -> str: diff --git a/optuna_dashboard/ts/action.ts b/optuna_dashboard/ts/action.ts index da87dc3f..335a815c 100644 --- a/optuna_dashboard/ts/action.ts +++ b/optuna_dashboard/ts/action.ts @@ -3,6 +3,7 @@ import { useSnackbar } from "notistack" import { getStudyDetailAPI, getStudySummariesAPI, + getParamImportances, createNewStudyAPI, deleteStudyAPI, saveNoteAPI, @@ -11,6 +12,7 @@ import { graphVisibilityState, studyDetailsState, studySummariesState, + paramImportanceState, } from "./state" const localStorageGraphVisibility = "graphVisibility" @@ -23,6 +25,8 @@ export const actionCreator = () => { useRecoilState(studyDetailsState) const [graphVisibility, setGraphVisibility] = useRecoilState(graphVisibilityState) + const [paramImportance, setParamImportance] = + useRecoilState(paramImportanceState) const setStudyDetailState = (studyId: number, study: StudyDetail) => { const newVal = Object.assign({}, studyDetails) @@ -30,6 +34,15 @@ export const actionCreator = () => { setStudyDetails(newVal) } + const setStudyParamImportanceState = ( + studyId: number, + importance: ParamImportance[][] + ) => { + const newVal = Object.assign({}, paramImportance) + newVal[studyId] = importance + setParamImportance(newVal) + } + const updateStudySummaries = (successMsg?: string) => { getStudySummariesAPI() .then((studySummaries: StudySummary[]) => { @@ -77,6 +90,22 @@ export const actionCreator = () => { }) } + const updateParamImportance = (studyId: number) => { + getParamImportances(studyId) + .then((importance) => { + setStudyParamImportanceState(studyId, importance) + }) + .catch((err) => { + const reason = err.response?.data.reason + enqueueSnackbar( + `Failed to load hyperparameter importance (reason=${reason})`, + { + variant: "error", + } + ) + }) + } + const createNewStudy = (studyName: string, directions: StudyDirection[]) => { createNewStudyAPI(studyName, directions) .then((study_summary) => { @@ -157,6 +186,7 @@ export const actionCreator = () => { return { updateStudyDetail, updateStudySummaries, + updateParamImportance, createNewStudy, deleteStudy, getGraphVisibility, diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index d7812a56..405e18f6 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -195,24 +195,15 @@ export const saveNoteAPI = ( } interface ParamImportancesResponse { - target_name: string - param_importances: ParamImportance[] + param_importances: ParamImportance[][] } export const getParamImportances = ( - studyId: number, - objectiveId = 0 -): Promise => { + studyId: number +): Promise => { return axiosInstance - .get( - `/api/studies/${studyId}/param_importances`, - { - params: { - objective_id: objectiveId, - }, - } - ) + .get(`/api/studies/${studyId}/param_importances`) .then((res) => { - return res.data + return res.data.param_importances }) } diff --git a/optuna_dashboard/ts/components/GraphHyperparameterImportances.tsx b/optuna_dashboard/ts/components/GraphHyperparameterImportances.tsx index fa7451e3..9cfa08b1 100644 --- a/optuna_dashboard/ts/components/GraphHyperparameterImportances.tsx +++ b/optuna_dashboard/ts/components/GraphHyperparameterImportances.tsx @@ -1,5 +1,6 @@ import * as plotly from "plotly.js-dist-min" import React, { FC, useEffect, useState } from "react" +import { useRecoilValue } from "recoil" import { Grid, FormControl, @@ -12,49 +13,43 @@ import { Box, } from "@mui/material" -import { getParamImportances } from "../apiClient" import { plotlyDarkTemplate } from "./PlotlyDarkMode" -import { useSnackbar } from "notistack" +import { actionCreator } from "../action" +import { paramImportanceState } from "../state" const plotDomId = "graph-hyperparameter-importances" +const useParamImportanceValue = ( + studyId: number +): ParamImportance[][] | null => { + const studyParamImportance = + useRecoilValue(paramImportanceState) + return studyParamImportance[studyId] || null +} + export const GraphHyperparameterImportances: FC<{ study: StudyDetail | null studyId: number }> = ({ study = null, studyId }) => { const theme = useTheme() + const action = actionCreator() + const importances = useParamImportanceValue(studyId) const [objectiveId, setObjectiveId] = useState(0) const numCompletedTrials = study?.trials.filter((t) => t.state === "Complete").length || 0 - const [importances, setImportances] = useState(null) - const { enqueueSnackbar } = useSnackbar() const handleObjectiveChange = (event: SelectChangeEvent) => { setObjectiveId(event.target.value as number) } useEffect(() => { - if (numCompletedTrials > 0) { - getParamImportances(studyId, objectiveId) - .then((p) => { - setImportances(p) - }) - .catch((err) => { - const reason = err.response?.data.reason - enqueueSnackbar( - `Failed to load hyperparameter importance (reason=${reason})`, - { - variant: "error", - } - ) - }) - } - }, [numCompletedTrials, objectiveId, theme.palette.mode]) + action.updateParamImportance(studyId) + }, [numCompletedTrials]) useEffect(() => { - if (importances !== null) { - plotParamImportances(importances, theme.palette.mode) + if (importances !== null && importances.length > objectiveId) { + plotParamImportances(importances[objectiveId], theme.palette.mode) } - }, [importances, theme.palette.mode]) + }, [importances, objectiveId, theme.palette.mode]) return ( @@ -88,25 +83,20 @@ export const GraphHyperparameterImportances: FC<{ ) } -const plotParamImportances = ( - paramsImportanceData: ParamImportances, - mode: string -) => { +const plotParamImportances = (importance: ParamImportance[], mode: string) => { if (document.getElementById(plotDomId) === null) { return } - const param_importances = [ - ...paramsImportanceData.param_importances, - ].reverse() - const importance_values = param_importances.map((p) => p.importance) - const param_names = param_importances.map((p) => p.name) - const param_hover_templates = param_importances.map( + const reversed = [...importance].reverse() + const importance_values = reversed.map((p) => p.importance) + const param_names = reversed.map((p) => p.name) + const param_hover_templates = reversed.map( (p) => `${p.name} (${p.distribution}): ${p.importance} ` ) const layout: Partial = { xaxis: { - title: `Importance for ${paramsImportanceData.target_name}`, + title: `Importance for the Objective Value`, }, yaxis: { title: "Hyperparameter", diff --git a/optuna_dashboard/ts/state.ts b/optuna_dashboard/ts/state.ts index e0ea82bb..5dda1de3 100644 --- a/optuna_dashboard/ts/state.ts +++ b/optuna_dashboard/ts/state.ts @@ -10,6 +10,11 @@ export const studyDetailsState = atom({ default: {}, }) +export const paramImportanceState = atom({ + key: "paramImportance", + default: {}, +}) + export const graphVisibilityState = atom({ key: "graphVisibility", default: { diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index 1cb85277..c5a64b57 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -109,7 +109,6 @@ declare interface StudyDetails { [study_id: string]: StudyDetail } -declare interface ParamImportances { - target_name: string - param_importances: ParamImportance[] +declare interface StudyParamImportance { + [study_id: string]: ParamImportance[][] } From 977a43b54a78f56b8c5f0c713838326193d0e863 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 1 Jan 2023 12:34:23 +0900 Subject: [PATCH 02/21] Improve importance plot on beta ui --- .../GraphHyperparameterImportances.tsx | 100 ++++++++++++++++-- .../ts/components/StudyDetailBeta.tsx | 14 +-- optuna_dashboard/ts/state.ts | 28 ++++- 3 files changed, 125 insertions(+), 17 deletions(-) diff --git a/optuna_dashboard/ts/components/GraphHyperparameterImportances.tsx b/optuna_dashboard/ts/components/GraphHyperparameterImportances.tsx index 9cfa08b1..d1eee8ae 100644 --- a/optuna_dashboard/ts/components/GraphHyperparameterImportances.tsx +++ b/optuna_dashboard/ts/components/GraphHyperparameterImportances.tsx @@ -1,6 +1,5 @@ import * as plotly from "plotly.js-dist-min" import React, { FC, useEffect, useState } from "react" -import { useRecoilValue } from "recoil" import { Grid, FormControl, @@ -11,19 +10,106 @@ import { SelectChangeEvent, useTheme, Box, + Card, + CardContent, } from "@mui/material" import { plotlyDarkTemplate } from "./PlotlyDarkMode" import { actionCreator } from "../action" -import { paramImportanceState } from "../state" +import { useParamImportanceValue, useStudyDirections } from "../state" const plotDomId = "graph-hyperparameter-importances" +const getPlotDomId = (objectiveId: number) => `graph-importance-${objectiveId}` -const useParamImportanceValue = ( +export const GraphHyperparameterImportanceBeta: FC<{ studyId: number -): ParamImportance[][] | null => { - const studyParamImportance = - useRecoilValue(paramImportanceState) - return studyParamImportance[studyId] || null + study: StudyDetail | null +}> = ({ studyId, study = null }) => { + const theme = useTheme() + const action = actionCreator() + const importances = useParamImportanceValue(studyId) + const numCompletedTrials = + study?.trials.filter((t) => t.state === "Complete").length || 0 + const nObjectives = useStudyDirections(studyId)?.length + + useEffect(() => { + action.updateParamImportance(studyId) + }, [numCompletedTrials]) + + useEffect(() => { + if (importances !== null && nObjectives === importances.length) { + plotParamImportancesBeta(importances, theme.palette.mode) + } + }, [nObjectives, importances, theme.palette.mode]) + + return ( + + {Array.from({ length: nObjectives || 1 }, (_, i) => ( + + + + + + + + ))} + + ) +} + +const plotParamImportancesBeta = ( + importances: ParamImportance[][], + mode: string +) => { + const getLayout = (title: string): Partial => ({ + xaxis: { + title: title, + }, + yaxis: { + title: "Hyperparameter", + automargin: true, + }, + margin: { + l: 50, + t: 0, + r: 50, + b: 50, + }, + showlegend: false, + template: mode === "dark" ? plotlyDarkTemplate : {}, + }) + + importances.forEach((importance, objectiveId) => { + if (document.getElementById(getPlotDomId(objectiveId)) === null) { + return + } + + const reversed = [...importance].reverse() + const importance_values = reversed.map((p) => p.importance) + const param_names = reversed.map((p) => p.name) + const param_hover_templates = reversed.map( + (p) => `${p.name} (${p.distribution}): ${p.importance} ` + ) + let title = `Importance for the Objective Value` + if (importance.length > 1) { + title = `Importance for the Objective ${objectiveId}` + } + const layout = getLayout(title) + const plotData: Partial[] = [ + { + type: "bar", + orientation: "h", + x: importance_values, + y: param_names, + text: importance_values.map((v) => String(v.toFixed(2))), + textposition: "outside", + hovertemplate: param_hover_templates, + marker: { + color: "rgb(66,146,198)", + }, + }, + ] + plotly.react(getPlotDomId(objectiveId), plotData, layout) + }) } export const GraphHyperparameterImportances: FC<{ diff --git a/optuna_dashboard/ts/components/StudyDetailBeta.tsx b/optuna_dashboard/ts/components/StudyDetailBeta.tsx index e4bac275..58f9c1ec 100644 --- a/optuna_dashboard/ts/components/StudyDetailBeta.tsx +++ b/optuna_dashboard/ts/components/StudyDetailBeta.tsx @@ -26,7 +26,7 @@ import { TrialTable } from "./TrialTable" import { AppDrawer } from "./AppDrawer" import { GraphParallelCoordinate } from "./GraphParallelCoordinate" import { Contour } from "./GraphContour" -import { GraphHyperparameterImportances } from "./GraphHyperparameterImportances" +import { GraphHyperparameterImportanceBeta } from "./GraphHyperparameterImportances" import { GraphSlice } from "./GraphSlice" import { GraphParetoFront } from "./GraphParetoFront" import { DataGrid, DataGridColumn } from "./DataGrid" @@ -188,14 +188,10 @@ export const StudyDetailBeta: FC<{ Hyperparameter Importance - - - - - + Hyperparameter Relationships diff --git a/optuna_dashboard/ts/state.ts b/optuna_dashboard/ts/state.ts index 5dda1de3..8dcad426 100644 --- a/optuna_dashboard/ts/state.ts +++ b/optuna_dashboard/ts/state.ts @@ -1,4 +1,4 @@ -import { atom } from "recoil" +import { atom, useRecoilValue } from "recoil" export const studySummariesState = atom({ key: "studySummaries", @@ -33,3 +33,29 @@ export const reloadIntervalState = atom({ key: "reloadInterval", default: 10, }) + +export const useStudyDetailValue = (studyId: number): StudyDetail | null => { + const studyDetails = useRecoilValue(studyDetailsState) + return studyDetails[studyId] || null +} + +export const useStudySummaryValue = (studyId: number): StudySummary | null => { + const studySummaries = useRecoilValue(studySummariesState) + return studySummaries.find((s) => s.study_id == studyId) || null +} + +export const useParamImportanceValue = ( + studyId: number +): ParamImportance[][] | null => { + const studyParamImportance = + useRecoilValue(paramImportanceState) + return studyParamImportance[studyId] || null +} + +export const useStudyDirections = ( + studyId: number +): StudyDirection[] | null => { + const studyDetail = useStudyDetailValue(studyId) + const studySummary = useStudySummaryValue(studyId) + return studyDetail?.directions || studySummary?.directions || null +} From 3d7dbcf31410fc373246e92f72e7ad767e949971 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 1 Jan 2023 12:45:45 +0900 Subject: [PATCH 03/21] Make drawer open state global --- optuna_dashboard/ts/components/AppDrawer.tsx | 4 ++-- optuna_dashboard/ts/components/StudyDetailBeta.tsx | 14 +++++++------- optuna_dashboard/ts/components/StudyListBeta.tsx | 7 ++----- optuna_dashboard/ts/state.ts | 5 +++++ 4 files changed, 16 insertions(+), 14 deletions(-) diff --git a/optuna_dashboard/ts/components/AppDrawer.tsx b/optuna_dashboard/ts/components/AppDrawer.tsx index 44154bae..99a0fee3 100644 --- a/optuna_dashboard/ts/components/AppDrawer.tsx +++ b/optuna_dashboard/ts/components/AppDrawer.tsx @@ -14,7 +14,7 @@ import ListItem from "@mui/material/ListItem" import ListItemButton from "@mui/material/ListItemButton" import ListItemIcon from "@mui/material/ListItemIcon" import ListItemText from "@mui/material/ListItemText" -import { reloadIntervalState } from "../state" +import { drawerOpenState, reloadIntervalState } from "../state" import { Link } from "react-router-dom" import AutoGraphIcon from "@mui/icons-material/AutoGraph" import SyncIcon from "@mui/icons-material/Sync" @@ -109,7 +109,7 @@ export const AppDrawer: FC<{ children?: React.ReactNode }> = ({ studyId, toggleColorMode, page, toolbar, children }) => { const theme = useTheme() - const [open, setOpen] = React.useState(false) + const [open, setOpen] = useRecoilState(drawerOpenState) const [reloadInterval, updateReloadInterval] = useRecoilState(reloadIntervalState) diff --git a/optuna_dashboard/ts/components/StudyDetailBeta.tsx b/optuna_dashboard/ts/components/StudyDetailBeta.tsx index 58f9c1ec..f081c80c 100644 --- a/optuna_dashboard/ts/components/StudyDetailBeta.tsx +++ b/optuna_dashboard/ts/components/StudyDetailBeta.tsx @@ -92,6 +92,13 @@ export const StudyDetailBeta: FC<{ if (page === "history") { content = ( + {directions !== null && directions.length > 1 ? ( + + + + + + ) : null} ) : null} - {directions !== null && directions.length > 1 ? ( - - - - - - ) : null} diff --git a/optuna_dashboard/ts/components/StudyListBeta.tsx b/optuna_dashboard/ts/components/StudyListBeta.tsx index 08977813..22efab9e 100644 --- a/optuna_dashboard/ts/components/StudyListBeta.tsx +++ b/optuna_dashboard/ts/components/StudyListBeta.tsx @@ -19,6 +19,7 @@ import { } from "@mui/material" import { Delete, Refresh, Search } from "@mui/icons-material" import SortIcon from "@mui/icons-material/Sort" +import HomeIcon from "@mui/icons-material/Home" import AddBoxIcon from "@mui/icons-material/AddBox" import { actionCreator } from "../action" @@ -101,11 +102,7 @@ export const StudyListBeta: FC<{ ) - const toolbar = ( - - Optuna Dashboard (Beta ver.) - - ) + const toolbar = return ( diff --git a/optuna_dashboard/ts/state.ts b/optuna_dashboard/ts/state.ts index 8dcad426..60ce15f0 100644 --- a/optuna_dashboard/ts/state.ts +++ b/optuna_dashboard/ts/state.ts @@ -34,6 +34,11 @@ export const reloadIntervalState = atom({ default: 10, }) +export const drawerOpenState = atom({ + key: "drawerOpen", + default: false, +}) + export const useStudyDetailValue = (studyId: number): StudyDetail | null => { const studyDetails = useRecoilValue(studyDetailsState) return studyDetails[studyId] || null From af17c58e4b4b3042d74d7b83b44142836594c66f Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 1 Jan 2023 13:17:30 +0900 Subject: [PATCH 04/21] Add refactor changes --- optuna_dashboard/ts/components/AppDrawer.tsx | 2 +- .../ts/components/StudyDetailBeta.tsx | 25 ++++++------------- optuna_dashboard/ts/state.ts | 6 +++++ optuna_dashboard/ts/types/index.d.ts | 2 ++ 4 files changed, 16 insertions(+), 19 deletions(-) diff --git a/optuna_dashboard/ts/components/AppDrawer.tsx b/optuna_dashboard/ts/components/AppDrawer.tsx index 99a0fee3..5af537b4 100644 --- a/optuna_dashboard/ts/components/AppDrawer.tsx +++ b/optuna_dashboard/ts/components/AppDrawer.tsx @@ -104,7 +104,7 @@ const Drawer = styled(MuiDrawer, { export const AppDrawer: FC<{ studyId?: number toggleColorMode: () => void - page?: "history" | "analytics" | "trials" | "note" + page?: PageId toolbar: React.ReactNode children?: React.ReactNode }> = ({ studyId, toggleColorMode, page, toolbar, children }) => { diff --git a/optuna_dashboard/ts/components/StudyDetailBeta.tsx b/optuna_dashboard/ts/components/StudyDetailBeta.tsx index f081c80c..c66b0b79 100644 --- a/optuna_dashboard/ts/components/StudyDetailBeta.tsx +++ b/optuna_dashboard/ts/components/StudyDetailBeta.tsx @@ -19,8 +19,10 @@ import { Note } from "./Note" import { actionCreator } from "../action" import { reloadIntervalState, - studyDetailsState, - studySummariesState, + useStudyDetailValue, + useStudyDirections, + useStudyName, + useStudySummaryValue, } from "../state" import { TrialTable } from "./TrialTable" import { AppDrawer } from "./AppDrawer" @@ -37,18 +39,6 @@ interface ParamTypes { studyId: string } -type PageId = "history" | "analytics" | "trials" | "note" - -const useStudyDetailValue = (studyId: number): StudyDetail | null => { - const studyDetails = useRecoilValue(studyDetailsState) - return studyDetails[studyId] || null -} - -const useStudySummaryValue = (studyId: number): StudySummary | null => { - const studySummaries = useRecoilValue(studySummariesState) - return studySummaries.find((s) => s.study_id == studyId) || null -} - export const StudyDetailBeta: FC<{ toggleColorMode: () => void page: PageId @@ -60,13 +50,12 @@ export const StudyDetailBeta: FC<{ const studyDetail = useStudyDetailValue(studyIdNumber) const reloadInterval = useRecoilValue(reloadIntervalState) const studySummary = useStudySummaryValue(studyIdNumber) - const directions = studyDetail?.directions || studySummary?.directions || null + const directions = useStudyDirections(studyIdNumber) + const studyName = useStudyName(studyIdNumber) const userAttrs = studySummary?.user_attrs || [] const title = - studyDetail !== null || studySummary !== null - ? `${studyDetail?.name || studySummary?.study_name} (id=${studyId})` - : `Study #${studyId}` + studyName !== null ? `${studyName} (id=${studyId})` : `Study #${studyId}` useEffect(() => { action.updateStudyDetail(studyIdNumber) diff --git a/optuna_dashboard/ts/state.ts b/optuna_dashboard/ts/state.ts index 60ce15f0..d298bf6b 100644 --- a/optuna_dashboard/ts/state.ts +++ b/optuna_dashboard/ts/state.ts @@ -64,3 +64,9 @@ export const useStudyDirections = ( const studySummary = useStudySummaryValue(studyId) return studyDetail?.directions || studySummary?.directions || null } + +export const useStudyName = (studyId: number): string | null => { + const studyDetail = useStudyDetailValue(studyId) + const studySummary = useStudySummaryValue(studyId) + return studyDetail?.name || studySummary?.study_name || null +} diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index c5a64b57..05d5fbd2 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -21,6 +21,8 @@ type Distribution = | "IntLogUniformDistribution" | "CategoricalDistribution" +type PageId = "history" | "analytics" | "trials" | "note" + type GraphVisibility = { history: boolean paretoFront: boolean From 2bca3d5feefde7323cde9615ae2b1047de970684 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 1 Jan 2023 13:36:27 +0900 Subject: [PATCH 05/21] Support best_trials --- .../ts/components/StudyDetailBeta.tsx | 46 +++++++++++++++---- 1 file changed, 36 insertions(+), 10 deletions(-) diff --git a/optuna_dashboard/ts/components/StudyDetailBeta.tsx b/optuna_dashboard/ts/components/StudyDetailBeta.tsx index c66b0b79..a9de119b 100644 --- a/optuna_dashboard/ts/components/StudyDetailBeta.tsx +++ b/optuna_dashboard/ts/components/StudyDetailBeta.tsx @@ -120,29 +120,55 @@ export const StudyDetailBeta: FC<{ Best Trial {studyDetail.best_trials[0].values} - - {studyDetail.best_trials[0].params.map((param) => ( - - {param.name} {param.value} - - ))} - + + Params = [ + {studyDetail.best_trials[0].params + .map((p) => `${p.name}: ${p.value}`) + .join(", ")} + ] + )} - {studyDetail !== null && studyDetail.best_trials.length > 1 && ( + {studyDetail !== null && studyDetail.directions.length > 1 && ( <> - Best Trials + Best Trials ({studyDetail.best_trials.length} trials) + {studyDetail.best_trials.map((trial, i) => ( + + + + Trial number={trial.number} (trial_id= + {trial.trial_id}) + + + Objective Values = [{trial.values?.join(", ")}] + + + Params = [ + {trial.params + .map((p) => `${p.name}: ${p.value}`) + .join(", ")} + ] + + + + ))} )} From f0851d565ef7340d7e78dd26bc7bd6eeff5eed76 Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 1 Jan 2023 14:05:21 +0900 Subject: [PATCH 06/21] Slightly improve StudyList --- optuna_dashboard/ts/components/StudyDetailBeta.tsx | 2 -- optuna_dashboard/ts/components/StudyListBeta.tsx | 2 +- 2 files changed, 1 insertion(+), 3 deletions(-) diff --git a/optuna_dashboard/ts/components/StudyDetailBeta.tsx b/optuna_dashboard/ts/components/StudyDetailBeta.tsx index a9de119b..26e56307 100644 --- a/optuna_dashboard/ts/components/StudyDetailBeta.tsx +++ b/optuna_dashboard/ts/components/StudyDetailBeta.tsx @@ -7,7 +7,6 @@ import { Box, Typography, useTheme, - ListItem, IconButton, } from "@mui/material" import Grid2 from "@mui/material/Unstable_Grid2" @@ -32,7 +31,6 @@ import { GraphHyperparameterImportanceBeta } from "./GraphHyperparameterImportan import { GraphSlice } from "./GraphSlice" import { GraphParetoFront } from "./GraphParetoFront" import { DataGrid, DataGridColumn } from "./DataGrid" -import List from "@mui/material/List" import { GraphIntermediateValues } from "./GraphIntermediateValues" interface ParamTypes { diff --git a/optuna_dashboard/ts/components/StudyListBeta.tsx b/optuna_dashboard/ts/components/StudyListBeta.tsx index 22efab9e..a05c8f43 100644 --- a/optuna_dashboard/ts/components/StudyListBeta.tsx +++ b/optuna_dashboard/ts/components/StudyListBeta.tsx @@ -178,7 +178,7 @@ export const StudyListBeta: FC<{ > - {study.study_id} {study.study_name} + {study.study_id}. {study.study_name} Date: Sun, 1 Jan 2023 14:47:22 +0900 Subject: [PATCH 07/21] Remove preventDefault --- optuna_dashboard/ts/components/GraphSlice.tsx | 1 - 1 file changed, 1 deletion(-) diff --git a/optuna_dashboard/ts/components/GraphSlice.tsx b/optuna_dashboard/ts/components/GraphSlice.tsx index 609e8aab..2a101237 100644 --- a/optuna_dashboard/ts/components/GraphSlice.tsx +++ b/optuna_dashboard/ts/components/GraphSlice.tsx @@ -61,7 +61,6 @@ export const GraphSlice: FC<{ } const handleLogYScaleChange = (e: ChangeEvent) => { - e.preventDefault() setLogYScale(!logYScale) } From 006c79be2ba1431c5cd90abc47b14864a4af4b9e Mon Sep 17 00:00:00 2001 From: c-bata Date: Sun, 1 Jan 2023 15:35:12 +0900 Subject: [PATCH 08/21] Slightly improve StudyList --- optuna_dashboard/ts/components/StudyListBeta.tsx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/ts/components/StudyListBeta.tsx b/optuna_dashboard/ts/components/StudyListBeta.tsx index a05c8f43..02175a19 100644 --- a/optuna_dashboard/ts/components/StudyListBeta.tsx +++ b/optuna_dashboard/ts/components/StudyListBeta.tsx @@ -150,7 +150,7 @@ export const StudyListBeta: FC<{ }} sx={{ marginRight: theme.spacing(2), minWidth: "120px" }} > - Refresh + Reload