From 1a7ac3ff3c828c9554d4b410eb2e70017305f47b Mon Sep 17 00:00:00 2001 From: Henry Cui Date: Sat, 27 Feb 2021 16:16:49 +0900 Subject: [PATCH 1/6] Add simple pareto front plot --- optuna_dashboard/static/apiClient.ts | 9 +++ .../static/components/GraphParetoFront.tsx | 60 +++++++++++++++++++ .../static/components/StudyDetail.tsx | 8 +++ optuna_dashboard/static/types/index.d.ts | 1 + 4 files changed, 78 insertions(+) create mode 100644 optuna_dashboard/static/components/GraphParetoFront.tsx diff --git a/optuna_dashboard/static/apiClient.ts b/optuna_dashboard/static/apiClient.ts index cd68c40a..befda07a 100644 --- a/optuna_dashboard/static/apiClient.ts +++ b/optuna_dashboard/static/apiClient.ts @@ -34,11 +34,17 @@ const convertTrialResponse = (res: TrialResponse): Trial => { } } +const convertTrialResponseList = (res: TrialResponse[]): Trial[] => { + return res.map((trial): Trial => convertTrialResponse(trial)) +} + + interface StudyDetailResponse { name: string datetime_start: string directions: StudyDirection[] best_trial?: TrialResponse + best_trials?: TrialResponse[] trials: TrialResponse[] } @@ -58,6 +64,9 @@ export const getStudyDetailAPI = (studyId: number): Promise => { best_trial: res.data.best_trial ? convertTrialResponse(res.data.best_trial) : undefined, + best_trials: res.data.best_trials + ? convertTrialResponseList(res.data.best_trials) + : undefined, trials: trials, } }) diff --git a/optuna_dashboard/static/components/GraphParetoFront.tsx b/optuna_dashboard/static/components/GraphParetoFront.tsx new file mode 100644 index 00000000..196e0a61 --- /dev/null +++ b/optuna_dashboard/static/components/GraphParetoFront.tsx @@ -0,0 +1,60 @@ +import * as plotly from "plotly.js-dist" +import React, { FC, useEffect } from "react" + +const plotDomId = "graph-pareto-front" + +export const GraphParetoFront: FC<{ + study: StudyDetail | null +}> = ({ study = null }) => { + + useEffect(() => { + if (study != null) { + plotParetoFront(study) + } + }, [study]) + + return
+} + +const plotParetoFront = (study: StudyDetail) => { + if (document.getElementById(plotDomId) === null) { + return + } + + if (study.directions.length != 2) { + return + } + + const layout: Partial = { + title: "Pareto-front plot", + margin: { + l: 50, + r: 50, + b: 0, + }, + } + + const trials: Trial[] = (study !== null) && study.best_trials ? study.best_trials : [] + console.log(study.best_trials) + console.log('length', trials.length) + if (trials.length === 0) { + plotly.react(plotDomId, [], layout) + return + } + + const pointColors = Array(trials.length).fill("blue") + + const plotData: Partial[] = [ + { + type: "scatter", + x: trials.map((t: Trial): number => t.values![0]), + y: trials.map((t: Trial): number => t.values![1]), + mode: "markers", + marker: { + color: pointColors, + }, + }, + ] + + plotly.react(plotDomId, plotData, layout) +} diff --git a/optuna_dashboard/static/components/StudyDetail.tsx b/optuna_dashboard/static/components/StudyDetail.tsx index d4068fec..85e4dc7e 100644 --- a/optuna_dashboard/static/components/StudyDetail.tsx +++ b/optuna_dashboard/static/components/StudyDetail.tsx @@ -23,6 +23,7 @@ import { GraphParallelCoordinate } from "./GraphParallelCoordinate" import { GraphIntermediateValues } from "./GraphIntermediateValues" import { GraphSlice } from "./GraphSlice" import { GraphHistory } from "./GraphHistory" +import { GraphParetoFront } from "./GraphParetoFront" import { actionCreator } from "../action" import { studyDetailsState } from "../state" @@ -203,6 +204,13 @@ export const StudyDetail: FC = () => { ) : null} + {studyDetail !== null && !isSingleObjectiveStudy(studyDetail) ? ( + + + + + + ) : null} diff --git a/optuna_dashboard/static/types/index.d.ts b/optuna_dashboard/static/types/index.d.ts index a534a02f..295a24c7 100644 --- a/optuna_dashboard/static/types/index.d.ts +++ b/optuna_dashboard/static/types/index.d.ts @@ -54,6 +54,7 @@ declare interface StudyDetail { directions: StudyDirection[] datetime_start: Date best_trial?: Trial + best_trials?: Trial[] trials: Trial[] } From 2a0b7e0d625fca3a242153edc646c6965bf4f260 Mon Sep 17 00:00:00 2001 From: Henry Cui Date: Sat, 27 Feb 2021 17:50:04 +0900 Subject: [PATCH 2/6] Enabled pareto front 2d plot --- optuna_dashboard/static/apiClient.ts | 9 ---- .../static/components/GraphParetoFront.tsx | 46 +++++++++++++++---- optuna_dashboard/static/types/index.d.ts | 1 - 3 files changed, 38 insertions(+), 18 deletions(-) diff --git a/optuna_dashboard/static/apiClient.ts b/optuna_dashboard/static/apiClient.ts index befda07a..cd68c40a 100644 --- a/optuna_dashboard/static/apiClient.ts +++ b/optuna_dashboard/static/apiClient.ts @@ -34,17 +34,11 @@ const convertTrialResponse = (res: TrialResponse): Trial => { } } -const convertTrialResponseList = (res: TrialResponse[]): Trial[] => { - return res.map((trial): Trial => convertTrialResponse(trial)) -} - - interface StudyDetailResponse { name: string datetime_start: string directions: StudyDirection[] best_trial?: TrialResponse - best_trials?: TrialResponse[] trials: TrialResponse[] } @@ -64,9 +58,6 @@ export const getStudyDetailAPI = (studyId: number): Promise => { best_trial: res.data.best_trial ? convertTrialResponse(res.data.best_trial) : undefined, - best_trials: res.data.best_trials - ? convertTrialResponseList(res.data.best_trials) - : undefined, trials: trials, } }) diff --git a/optuna_dashboard/static/components/GraphParetoFront.tsx b/optuna_dashboard/static/components/GraphParetoFront.tsx index 196e0a61..99bba3f3 100644 --- a/optuna_dashboard/static/components/GraphParetoFront.tsx +++ b/optuna_dashboard/static/components/GraphParetoFront.tsx @@ -21,7 +21,8 @@ const plotParetoFront = (study: StudyDetail) => { return } - if (study.directions.length != 2) { + const dim: number = study.directions.length + if (dim != 2) { return } @@ -34,21 +35,50 @@ const plotParetoFront = (study: StudyDetail) => { }, } - const trials: Trial[] = (study !== null) && study.best_trials ? study.best_trials : [] - console.log(study.best_trials) - console.log('length', trials.length) - if (trials.length === 0) { + const trials: Trial[] = study !== null ? study.trials : [] + const completedTrials = trials.filter( + (t) => t.state === "Complete" + ) + + if (completedTrials.length === 0) { plotly.react(plotDomId, [], layout) return } - const pointColors = Array(trials.length).fill("blue") + const normalizedValues: number[][] = [] + completedTrials.forEach((t) => { + if (t.values && t.values.length == dim) { + let values: number[] = t.values + values.forEach((v: number, i: number) => { + if (study.directions[i] === "maximize") { + values[i] = -v + } + }) + normalizedValues.push(values) + } + }) + + const pointColors: string[] = [] + normalizedValues.forEach((values0: number[], i: number) => { + let dominated: boolean = false + + dominated = normalizedValues.some((values1: number[], j: number) => { + if (i === j) { + return false + } + return values0.every((value0: number, k: number) => { + return value0 <= values1[k] + }) + }) + + if (dominated) { pointColors.push("blue") } else { pointColors.push("red") } + }) const plotData: Partial[] = [ { type: "scatter", - x: trials.map((t: Trial): number => t.values![0]), - y: trials.map((t: Trial): number => t.values![1]), + x: completedTrials.map((t: Trial): number => t.values![0]), + y: completedTrials.map((t: Trial): number => t.values![1]), mode: "markers", marker: { color: pointColors, diff --git a/optuna_dashboard/static/types/index.d.ts b/optuna_dashboard/static/types/index.d.ts index 295a24c7..a534a02f 100644 --- a/optuna_dashboard/static/types/index.d.ts +++ b/optuna_dashboard/static/types/index.d.ts @@ -54,7 +54,6 @@ declare interface StudyDetail { directions: StudyDirection[] datetime_start: Date best_trial?: Trial - best_trials?: Trial[] trials: Trial[] } From c4751307f1eb5bd8aff938d40c5c1489cd3d9fe0 Mon Sep 17 00:00:00 2001 From: Henry Cui Date: Sun, 28 Feb 2021 00:39:22 +0900 Subject: [PATCH 3/6] Debugged and improved --- .../static/components/GraphParetoFront.tsx | 34 ++++++++++++------- 1 file changed, 21 insertions(+), 13 deletions(-) diff --git a/optuna_dashboard/static/components/GraphParetoFront.tsx b/optuna_dashboard/static/components/GraphParetoFront.tsx index 99bba3f3..a1bbd040 100644 --- a/optuna_dashboard/static/components/GraphParetoFront.tsx +++ b/optuna_dashboard/static/components/GraphParetoFront.tsx @@ -25,7 +25,6 @@ const plotParetoFront = (study: StudyDetail) => { if (dim != 2) { return } - const layout: Partial = { title: "Pareto-front plot", margin: { @@ -45,29 +44,28 @@ const plotParetoFront = (study: StudyDetail) => { return } - const normalizedValues: number[][] = [] + let normalizedValues: number[][] = [] completedTrials.forEach((t) => { if (t.values && t.values.length == dim) { - let values: number[] = t.values - values.forEach((v: number, i: number) => { - if (study.directions[i] === "maximize") { - values[i] = -v + const trialValues = t.values.map( + (v: number, i: number) => { + return (study.directions[i] === "minimize") ? v : -v } - }) - normalizedValues.push(values) + ) + normalizedValues.push(trialValues) } }) const pointColors: string[] = [] normalizedValues.forEach((values0: number[], i: number) => { - let dominated: boolean = false + let dominated = false dominated = normalizedValues.some((values1: number[], j: number) => { if (i === j) { return false } return values0.every((value0: number, k: number) => { - return value0 <= values1[k] + return values1[k] <= value0 }) }) @@ -77,12 +75,22 @@ const plotParetoFront = (study: StudyDetail) => { const plotData: Partial[] = [ { type: "scatter", - x: completedTrials.map((t: Trial): number => t.values![0]), - y: completedTrials.map((t: Trial): number => t.values![1]), + x: completedTrials.map((t: Trial): number => { return t.values![0]}), + y: completedTrials.map((t: Trial): number => { return t.values![1]}), mode: "markers", + xaxis: "Objective 0", + yaxis: "Objective 1", marker: { - color: pointColors, + color: pointColors }, + text: completedTrials.map((t: Trial): string => { + return JSON.stringify({ + "number": t.number, + "values": t.values, + "params": t.params, + }, null, 2).replaceAll("\n", "
") + }), + hovertemplate: "%{text}", }, ] From 2ef70332b0a873ea4cd1f47b02f7b07d7ee2b667 Mon Sep 17 00:00:00 2001 From: Henry Cui Date: Sun, 28 Feb 2021 00:40:54 +0900 Subject: [PATCH 4/6] Change let to const for normalizedValues --- optuna_dashboard/static/components/GraphParetoFront.tsx | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/optuna_dashboard/static/components/GraphParetoFront.tsx b/optuna_dashboard/static/components/GraphParetoFront.tsx index a1bbd040..d8bcb3f3 100644 --- a/optuna_dashboard/static/components/GraphParetoFront.tsx +++ b/optuna_dashboard/static/components/GraphParetoFront.tsx @@ -44,7 +44,7 @@ const plotParetoFront = (study: StudyDetail) => { return } - let normalizedValues: number[][] = [] + const normalizedValues: number[][] = [] completedTrials.forEach((t) => { if (t.values && t.values.length == dim) { const trialValues = t.values.map( From 6f12ac71e38c9c069876134b8558c5ef967e0fef Mon Sep 17 00:00:00 2001 From: Henry Cui Date: Sun, 28 Feb 2021 15:55:29 +0900 Subject: [PATCH 5/6] Apply fmt --- .../static/components/GraphParetoFront.tsx | 61 +++++++++++-------- 1 file changed, 34 insertions(+), 27 deletions(-) diff --git a/optuna_dashboard/static/components/GraphParetoFront.tsx b/optuna_dashboard/static/components/GraphParetoFront.tsx index d8bcb3f3..f2a6f0cc 100644 --- a/optuna_dashboard/static/components/GraphParetoFront.tsx +++ b/optuna_dashboard/static/components/GraphParetoFront.tsx @@ -6,7 +6,6 @@ const plotDomId = "graph-pareto-front" export const GraphParetoFront: FC<{ study: StudyDetail | null }> = ({ study = null }) => { - useEffect(() => { if (study != null) { plotParetoFront(study) @@ -35,9 +34,7 @@ const plotParetoFront = (study: StudyDetail) => { } const trials: Trial[] = study !== null ? study.trials : [] - const completedTrials = trials.filter( - (t) => t.state === "Complete" - ) + const completedTrials = trials.filter((t) => t.state === "Complete") if (completedTrials.length === 0) { plotly.react(plotDomId, [], layout) @@ -47,11 +44,9 @@ const plotParetoFront = (study: StudyDetail) => { const normalizedValues: number[][] = [] completedTrials.forEach((t) => { if (t.values && t.values.length == dim) { - const trialValues = t.values.map( - (v: number, i: number) => { - return (study.directions[i] === "minimize") ? v : -v - } - ) + const trialValues = t.values.map((v: number, i: number) => { + return study.directions[i] === "minimize" ? v : -v + }) normalizedValues.push(trialValues) } }) @@ -69,28 +64,40 @@ const plotParetoFront = (study: StudyDetail) => { }) }) - if (dominated) { pointColors.push("blue") } else { pointColors.push("red") } + if (dominated) { + pointColors.push("blue") + } else { + pointColors.push("red") + } }) const plotData: Partial[] = [ { - type: "scatter", - x: completedTrials.map((t: Trial): number => { return t.values![0]}), - y: completedTrials.map((t: Trial): number => { return t.values![1]}), - mode: "markers", - xaxis: "Objective 0", - yaxis: "Objective 1", - marker: { - color: pointColors - }, - text: completedTrials.map((t: Trial): string => { - return JSON.stringify({ - "number": t.number, - "values": t.values, - "params": t.params, - }, null, 2).replaceAll("\n", "
") - }), - hovertemplate: "%{text}", + type: "scatter", + x: completedTrials.map((t: Trial): number => { + return t.values![0] + }), + y: completedTrials.map((t: Trial): number => { + return t.values![1] + }), + mode: "markers", + xaxis: "Objective 0", + yaxis: "Objective 1", + marker: { + color: pointColors, + }, + text: completedTrials.map((t: Trial): string => { + return JSON.stringify( + { + number: t.number, + values: t.values, + params: t.params, + }, + null, + 2 + ).replaceAll("\n", "
") + }), + hovertemplate: "%{text}", }, ] From 97422c63fcf2a8c617f076cf67885ede58f6e3b4 Mon Sep 17 00:00:00 2001 From: Henry Cui Date: Sat, 13 Mar 2021 15:51:23 +0900 Subject: [PATCH 6/6] Make axis selectable --- .../static/components/GraphParetoFront.tsx | 88 +++++++++++++++++-- 1 file changed, 79 insertions(+), 9 deletions(-) diff --git a/optuna_dashboard/static/components/GraphParetoFront.tsx b/optuna_dashboard/static/components/GraphParetoFront.tsx index f2a6f0cc..e0d4491b 100644 --- a/optuna_dashboard/static/components/GraphParetoFront.tsx +++ b/optuna_dashboard/static/components/GraphParetoFront.tsx @@ -1,21 +1,91 @@ import * as plotly from "plotly.js-dist" -import React, { FC, useEffect } from "react" +import React, { FC, useEffect, useState } from "react" +import { + Grid, + FormControl, + FormLabel, + MenuItem, + Select, +} from "@material-ui/core" +import { createStyles, makeStyles, Theme } from "@material-ui/core/styles" const plotDomId = "graph-pareto-front" +const useStyles = makeStyles((theme: Theme) => + createStyles({ + formControl: { + marginBottom: theme.spacing(2), + marginRight: theme.spacing(5), + marginTop: theme.spacing(10), + }, + }) +) + export const GraphParetoFront: FC<{ study: StudyDetail | null }> = ({ study = null }) => { + const classes = useStyles() + const [objectiveXId, setObjectiveXId] = useState(0) + const [objectiveYId, setObjectiveYId] = useState(1) + + const handleObjectiveXChange = ( + event: React.ChangeEvent<{ value: unknown }> + ) => { + setObjectiveXId(event.target.value as number) + } + + const handleObjectiveYChange = ( + event: React.ChangeEvent<{ value: unknown }> + ) => { + setObjectiveYId(event.target.value as number) + } + useEffect(() => { if (study != null) { - plotParetoFront(study) + plotParetoFront(study, objectiveXId, objectiveYId) } - }, [study]) + }, [study, objectiveXId, objectiveYId]) - return
+ return ( + + {study !== null && study.directions.length !== 1 ? ( + + + + Objective X ID: + + + + Objective Y ID: + + + + + ) : null} + +
+ + + ) } -const plotParetoFront = (study: StudyDetail) => { +const plotParetoFront = ( + study: StudyDetail, + objectiveXId: number, + objectiveYId: number +) => { if (document.getElementById(plotDomId) === null) { return } @@ -75,14 +145,14 @@ const plotParetoFront = (study: StudyDetail) => { { type: "scatter", x: completedTrials.map((t: Trial): number => { - return t.values![0] + return t.values![objectiveXId] }), y: completedTrials.map((t: Trial): number => { - return t.values![1] + return t.values![objectiveYId] }), mode: "markers", - xaxis: "Objective 0", - yaxis: "Objective 1", + xaxis: "Objective X", + yaxis: "Objective Y", marker: { color: pointColors, },