diff --git a/optuna_dashboard/ts/components/GraphHistory.tsx b/optuna_dashboard/ts/components/GraphHistory.tsx index 8a7fe3a9..42536e29 100644 --- a/optuna_dashboard/ts/components/GraphHistory.tsx +++ b/optuna_dashboard/ts/components/GraphHistory.tsx @@ -1,373 +1,32 @@ -import { - Box, - FormControl, - FormControlLabel, - FormLabel, - Grid, - MenuItem, - Radio, - RadioGroup, - Select, - SelectChangeEvent, - Slider, - Typography, - useTheme, -} from "@mui/material" -import { - Target, - useFilteredTrialsFromStudies, - useObjectiveAndUserAttrTargetsFromStudies, -} from "@optuna/react" -import * as Optuna from "@optuna/types" -import * as plotly from "plotly.js-dist-min" -import React, { ChangeEvent, FC, useEffect, useState } from "react" +import { useTheme } from "@mui/material" +import { PlotHistory } from "@optuna/react" +import React, { FC } from "react" import { useNavigate } from "react-router-dom" import { StudyDetail } from "ts/types/optuna" import { useConstants } from "../constantsProvider" import { usePlotlyColorTheme } from "../state" -const plotDomId = "graph-history" - -interface HistoryPlotInfo { - study_name: string - trials: Optuna.Trial[] - directions: Optuna.StudyDirection[] - metric_names?: string[] -} - export const GraphHistory: FC<{ studies: StudyDetail[] logScale: boolean includePruned: boolean }> = ({ studies, logScale, includePruned }) => { const { url_prefix } = useConstants() - const theme = useTheme() const colorTheme = usePlotlyColorTheme(theme.palette.mode) + const linkURL = (studyId: number, trialNumber: number) => { + return url_prefix + `/studies/${studyId}/trials?numbers=${trialNumber}` + } const navigate = useNavigate() - const [xAxis, setXAxis] = useState< - "number" | "datetime_start" | "datetime_complete" - >("number") - const [markerSize, setMarkerSize] = useState(5) - - const [targets, selected, setTarget] = - useObjectiveAndUserAttrTargetsFromStudies(studies) - - const trials = useFilteredTrialsFromStudies( - studies, - [selected], - !includePruned - ) - const historyPlotInfos = studies.map((study, index) => { - const h: HistoryPlotInfo = { - study_name: study?.name, - trials: trials[index], - directions: study?.directions, - metric_names: study?.metric_names, - } - return h - }) - - useEffect(() => { - plotHistory( - historyPlotInfos, - selected, - xAxis, - logScale, - theme.palette.mode, - colorTheme, - markerSize - ) - const element = document.getElementById(plotDomId) - if (element !== null && studies.length >= 1) { - // @ts-ignore - element.on("plotly_click", (data) => { - if (data.points[0].data.mode !== "lines") { - let studyId = 1 - if (data.points[0].data.name.includes("Infeasible Trial of")) { - const studyInfo: { id: number; name: string }[] = [] - studies.forEach((study) => { - studyInfo.push({ id: study.id, name: study.name }) - }) - const dataPointStudyName = data.points[0].data.name.replace( - "Infeasible Trial of ", - "" - ) - const targetId = studyInfo.find( - (s) => s.name === dataPointStudyName - )?.id - if (targetId !== undefined) { - studyId = targetId - } - } else { - studyId = studies[Math.floor(data.points[0].curveNumber / 2)].id - } - navigate( - url_prefix + - `/studies/${studyId}/trials?numbers=${data.points[0].x}` - ) - } - }) - return () => { - // @ts-ignore - element.removeAllListeners("plotly_click") - } - } - }, [ - studies, - selected, - logScale, - xAxis, - theme.palette.mode, - colorTheme, - markerSize, - ]) - - const handleObjectiveChange = (event: SelectChangeEvent) => { - setTarget(event.target.value) - } - - const handleXAxisChange = (e: ChangeEvent) => { - if (e.target.value === "number") { - setXAxis("number") - } else if (e.target.value === "datetime_start") { - setXAxis("datetime_start") - } else if (e.target.value === "datetime_complete") { - setXAxis("datetime_complete") - } - } return ( - - - - History - - {studies[0] !== null && targets.length >= 2 ? ( - - y Axis - - - ) : null} - - X-axis: - - } - label="Number" - /> - } - label="Datetime start" - /> - } - label="Datetime complete" - /> - - - - Marker size: - { - // @ts-ignore - setMarkerSize(e.target.value as number) - }} - /> - - - - - - + ) } - -const plotHistory = ( - historyPlotInfos: HistoryPlotInfo[], - target: Target, - xAxis: "number" | "datetime_start" | "datetime_complete", - logScale: boolean, - mode: string, - colorTheme: Partial, - markerSize: number -) => { - if (document.getElementById(plotDomId) === null) { - return - } - if (historyPlotInfos.length === 0) { - plotly.react(plotDomId, [], { - template: colorTheme, - }) - return - } - - const layout: Partial = { - margin: { - l: 50, - t: 0, - r: 50, - b: 0, - }, - yaxis: { - title: target.toLabel(historyPlotInfos[0].metric_names), - type: logScale ? "log" : "linear", - }, - xaxis: { - title: xAxis === "number" ? "Trial" : "Time", - type: xAxis === "number" ? "linear" : "date", - }, - showlegend: historyPlotInfos.length === 1 ? false : true, - template: colorTheme, - legend: { - x: 1.0, - y: 0.95, - }, - } - - const getAxisX = (trial: Optuna.Trial): number | Date => { - return xAxis === "number" - ? trial.number - : xAxis === "datetime_start" - ? trial.datetime_start ?? new Date() - : trial.datetime_complete ?? new Date() - } - - const plotData: Partial[] = [] - const infeasiblePlotData: Partial[] = [] - historyPlotInfos.forEach((h) => { - const feasibleTrials: Optuna.Trial[] = [] - const infeasibleTrials: Optuna.Trial[] = [] - h.trials.forEach((t) => { - if (t.constraints.every((c) => c <= 0)) { - feasibleTrials.push(t) - } else { - infeasibleTrials.push(t) - } - }) - plotData.push({ - x: feasibleTrials.map(getAxisX), - y: feasibleTrials.map( - (t: Optuna.Trial): number => target.getTargetValue(t) as number - ), - name: `${target.toLabel(h.metric_names)} of ${h.study_name}`, - marker: { - size: markerSize, - }, - mode: "markers", - type: "scatter", - }) - - const objectiveId = target.getObjectiveId() - if (objectiveId !== null) { - const xForLinePlot: (number | Date)[] = [] - const yForLinePlot: number[] = [] - let currentBest: number | null = null - for (let i = 0; i < feasibleTrials.length; i++) { - const t = feasibleTrials[i] - const value = target.getTargetValue(t) as number - if (t.state !== "Complete") { - continue - } else if (currentBest === null) { - currentBest = value - xForLinePlot.push(getAxisX(t)) - yForLinePlot.push(value) - } else if ( - h.directions[objectiveId] === "maximize" && - value > currentBest - ) { - const p = feasibleTrials[i - 1] - if (!xForLinePlot.includes(getAxisX(p))) { - xForLinePlot.push(getAxisX(p)) - yForLinePlot.push(currentBest) - } - currentBest = value - xForLinePlot.push(getAxisX(t)) - yForLinePlot.push(value) - } else if ( - h.directions[objectiveId] === "minimize" && - value < currentBest - ) { - const p = feasibleTrials[i - 1] - if (!xForLinePlot.includes(getAxisX(p))) { - xForLinePlot.push(getAxisX(p)) - yForLinePlot.push(currentBest) - } - currentBest = value - xForLinePlot.push(getAxisX(t)) - yForLinePlot.push(value) - } - } - if (h.trials.length !== 0) { - xForLinePlot.push(getAxisX(h.trials[h.trials.length - 1])) - yForLinePlot.push(yForLinePlot[yForLinePlot.length - 1]) - } - plotData.push({ - x: xForLinePlot, - y: yForLinePlot, - name: `Best Value of ${h.study_name}`, - mode: "lines", - type: "scatter", - }) - } - infeasiblePlotData.push({ - x: infeasibleTrials.map(getAxisX), - y: infeasibleTrials.map( - (t: Optuna.Trial): number => target.getTargetValue(t) as number - ), - name: `Infeasible Trial of ${h.study_name}`, - marker: { - size: markerSize, - color: mode === "dark" ? "#666666" : "#cccccc", - }, - mode: "markers", - type: "scatter", - showlegend: false, - }) - }) - plotData.push(...infeasiblePlotData) - plotly.react(plotDomId, plotData, layout) -} diff --git a/standalone_app/src/components/StudyDetail.tsx b/standalone_app/src/components/StudyDetail.tsx index ec37a2a0..42715824 100644 --- a/standalone_app/src/components/StudyDetail.tsx +++ b/standalone_app/src/components/StudyDetail.tsx @@ -173,7 +173,7 @@ export const StudyDetail: FC<{ - + {!!study && } diff --git a/tslib/react/src/components/PlotHistory.stories.tsx b/tslib/react/src/components/PlotHistory.stories.tsx index f274249e..159d836b 100644 --- a/tslib/react/src/components/PlotHistory.stories.tsx +++ b/tslib/react/src/components/PlotHistory.stories.tsx @@ -14,12 +14,13 @@ const meta: Meta = { (Story, storyContext) => { const { study } = useMockStudy(storyContext.parameters?.studyId) if (!study) return

loading...

+ const studies = [study] return ( diff --git a/tslib/react/src/components/PlotHistory.tsx b/tslib/react/src/components/PlotHistory.tsx index 1ea68db3..294dc7e0 100644 --- a/tslib/react/src/components/PlotHistory.tsx +++ b/tslib/react/src/components/PlotHistory.tsx @@ -1,5 +1,4 @@ import { - Checkbox, FormControl, FormControlLabel, FormLabel, @@ -9,6 +8,7 @@ import { RadioGroup, Select, SelectChangeEvent, + Slider, Switch, Typography, useTheme, @@ -16,60 +16,150 @@ import { import * as Optuna from "@optuna/types" import * as plotly from "plotly.js-dist-min" import { ChangeEvent, FC, useEffect, useState } from "react" + +import { useGraphComponentState } from "../hooks/useGraphComponentState" +import { + Target, + useFilteredTrialsFromStudies, + useObjectiveAndUserAttrTargetsFromStudies, +} from "../utils/trialFilter" import { plotlyDarkTemplate } from "./PlotlyDarkMode" const plotDomId = "plot-history" +interface HistoryPlotInfo { + study_name: string + trials: Optuna.Trial[] + directions: Optuna.StudyDirection[] + metric_names?: string[] +} + export const PlotHistory: FC<{ - study: Optuna.Study | null -}> = ({ study = null }) => { + studies: Optuna.Study[] + logScale?: boolean + includePruned?: boolean + colorTheme?: Partial + linkURL?: (studyId: number, trialNumber: number) => string + // biome-ignore lint/suspicious/noExplicitAny: It will accept any routers of each library. + router?: any +}> = ({ studies, logScale, includePruned, colorTheme, linkURL, router }) => { + const { graphComponentState, notifyGraphDidRender } = useGraphComponentState() + const theme = useTheme() - const [xAxis, setXAxis] = useState("number") - const [objectiveId, setObjectiveId] = useState(0) - const [logScale, setLogScale] = useState(false) - const [filterCompleteTrial, setFilterCompleteTrial] = useState(false) - const [filterPrunedTrial, setFilterPrunedTrial] = useState(false) + const colorThemeUsed = + colorTheme ?? (theme.palette.mode === "dark" ? plotlyDarkTemplate : {}) - const handleObjectiveChange = (event: SelectChangeEvent) => { - setObjectiveId(event.target.value as number) - } + const [xAxis, setXAxis] = useState< + "number" | "datetime_start" | "datetime_complete" + >("number") - const handleXAxisChange = (e: ChangeEvent) => { - setXAxis(e.target.value) + const [logScaleInternal, setLogScaleInternal] = useState(false) + const [includePrunedInternal, setIncludePrunedInternal] = + useState(true) + + const [markerSize, setMarkerSize] = useState(5) + + const [targets, selected, setTarget] = + useObjectiveAndUserAttrTargetsFromStudies(studies) + + const trials = useFilteredTrialsFromStudies( + studies, + [selected], + includePruned === undefined ? !includePrunedInternal : !includePruned + ) + + const historyPlotInfos = studies.map((study, index) => { + const h: HistoryPlotInfo = { + study_name: study.name, + trials: trials[index], + directions: study.directions, + metric_names: study.metric_names, + } + return h + }) + + const handleObjectiveChange = (event: SelectChangeEvent) => { + setTarget(event.target.value) } const handleLogScaleChange = () => { - setLogScale(!logScale) + setLogScaleInternal(!logScaleInternal) } - const handleFilterCompleteChange = () => { - setFilterCompleteTrial(!filterCompleteTrial) + const handleIncludePrunedChange = () => { + setIncludePrunedInternal(!includePrunedInternal) } - const handleFilterPrunedChange = () => { - setFilterPrunedTrial(!filterPrunedTrial) + const handleXAxisChange = (e: ChangeEvent) => { + if (e.target.value === "number") { + setXAxis("number") + } else if (e.target.value === "datetime_start") { + setXAxis("datetime_start") + } else if (e.target.value === "datetime_complete") { + setXAxis("datetime_complete") + } } + // biome-ignore lint/correctness/useExhaustiveDependencies: useEffect(() => { - if (study !== null) { + if (graphComponentState !== "componentWillMount") { plotHistory( - study, - objectiveId, + historyPlotInfos, + selected, xAxis, - logScale, - filterCompleteTrial, - filterPrunedTrial, - theme.palette.mode - ) + logScale === undefined ? logScaleInternal : logScale, + colorThemeUsed, + markerSize + )?.then(notifyGraphDidRender) + + const element = document.getElementById(plotDomId) + if ( + element !== null && + studies.length >= 1 && + linkURL !== undefined && + router !== undefined + ) { + // @ts-ignore + element.on("plotly_click", (data) => { + if (data.points[0].data.mode !== "lines") { + let studyId = 1 + if (data.points[0].data.name.includes("Infeasible Trial of")) { + const studyInfo: { id: number; name: string }[] = [] + for (const study of studies) { + studyInfo.push({ id: study.id, name: study.name }) + } + const dataPointStudyName = data.points[0].data.name.replace( + "Infeasible Trial of ", + "" + ) + const targetId = studyInfo.find( + (s) => s.name === dataPointStudyName + )?.id + if (targetId !== undefined) { + studyId = targetId + } + } else { + studyId = studies[Math.floor(data.points[0].curveNumber / 2)].id + } + const trialNumber = data.points[0].x + router(linkURL(studyId, trialNumber)) + } + }) + return () => { + // @ts-ignore + element.removeAllListeners("plotly_click") + } + } } }, [ - study, - objectiveId, - logScale, + historyPlotInfos, + selected, xAxis, - filterPrunedTrial, - filterCompleteTrial, - theme.palette.mode, + logScale, + logScaleInternal, + colorThemeUsed, + markerSize, + graphComponentState, ]) return ( @@ -81,60 +171,30 @@ export const PlotHistory: FC<{ direction="column" sx={{ paddingRight: theme.spacing(2) }} > - + History - {study !== null && study.directions.length !== 1 ? ( + {studies[0] !== null && targets.length >= 2 ? ( - Objective ID: - + {targets.map((t) => ( + + {t.toLabel(studies[0].metric_names)} ))} ) : null} - - Log y scale: - - - - Filter state: - - } - label="Complete" - /> - - } - label="Pruned" - /> - + {logScale === undefined ? ( + + Log y scale: + + + ) : null} + {includePruned === undefined ? ( + + Include PRUNED trials: + + + ) : null} + + Marker size: + { + // @ts-ignore + setMarkerSize(e.target.value as number) + }} + /> +
@@ -171,32 +271,23 @@ export const PlotHistory: FC<{ ) } -const filterFunc = (trial: Optuna.Trial, objectiveId: number): boolean => { - if (trial.state !== "Complete" && trial.state !== "Pruned") { - return false - } - if (trial.values === undefined) { - return false - } - return ( - trial.values.length > objectiveId && - trial.values[objectiveId] !== Infinity && - trial.values[objectiveId] !== -Infinity - ) -} - const plotHistory = ( - study: Optuna.Study, - objectiveId: number, - xAxis: string, + historyPlotInfos: HistoryPlotInfo[], + target: Target, + xAxis: "number" | "datetime_start" | "datetime_complete", logScale: boolean, - filterCompleteTrial: boolean, - filterPrunedTrial: boolean, - mode: string + colorTheme: Partial, + markerSize: number ) => { if (document.getElementById(plotDomId) === null) { return } + if (historyPlotInfos.length === 0) { + plotly.react(plotDomId, [], { + template: colorTheme, + }) + return + } const layout: Partial = { margin: { @@ -206,27 +297,19 @@ const plotHistory = ( b: 0, }, yaxis: { - title: "Objective Value", + title: target.toLabel(historyPlotInfos[0].metric_names), type: logScale ? "log" : "linear", }, xaxis: { title: xAxis === "number" ? "Trial" : "Time", type: xAxis === "number" ? "linear" : "date", }, - showlegend: true, - template: mode === "dark" ? plotlyDarkTemplate : {}, - } - - let filteredTrials = study.trials.filter((t) => filterFunc(t, objectiveId)) - if (filterCompleteTrial) { - filteredTrials = filteredTrials.filter((t) => t.state !== "Complete") - } - if (filterPrunedTrial) { - filteredTrials = filteredTrials.filter((t) => t.state !== "Pruned") - } - if (filteredTrials.length === 0) { - plotly.react(plotDomId, [], layout) - return + showlegend: historyPlotInfos.length === 1 ? false : true, + template: colorTheme, + legend: { + x: 1.0, + y: 0.95, + }, } const getAxisX = (trial: Optuna.Trial): number | Date => { @@ -237,80 +320,100 @@ const plotHistory = ( : trial.datetime_complete ?? new Date() } - const getValue = ( - trial: Optuna.Trial, - objectiveId: number - ): number | null => { - if ( - objectiveId === null || - trial.values === undefined || - trial.values.length <= objectiveId - ) { - return null - } - const value = trial.values[objectiveId] - if (value === Infinity || value === -Infinity) { - return null - } - return value - } - - const xForLinePlot: (number | Date)[] = [] - const yForLinePlot: number[] = [] - let currentBest: number | null = null - for (let i = 0; i < filteredTrials.length; i++) { - const t = filteredTrials[i] - const v = getValue(t, objectiveId) as number - if (currentBest === null) { - currentBest = v - xForLinePlot.push(getAxisX(t)) - yForLinePlot.push(v) - } else if ( - study.directions[objectiveId] === "maximize" && - v > currentBest - ) { - const p = filteredTrials[i - 1] - if (!xForLinePlot.includes(getAxisX(p))) { - xForLinePlot.push(getAxisX(p)) - yForLinePlot.push(currentBest) + const plotData: Partial[] = [] + const infeasiblePlotData: Partial[] = [] + for (const h of historyPlotInfos) { + const feasibleTrials: Optuna.Trial[] = [] + const infeasibleTrials: Optuna.Trial[] = [] + for (const t of h.trials) { + if (t.constraints.every((c) => c <= 0)) { + feasibleTrials.push(t) + } else { + infeasibleTrials.push(t) } - currentBest = v - xForLinePlot.push(getAxisX(t)) - yForLinePlot.push(v) - } else if ( - study.directions[objectiveId] === "minimize" && - v < currentBest - ) { - const p = filteredTrials[i - 1] - if (!xForLinePlot.includes(getAxisX(p))) { - xForLinePlot.push(getAxisX(p)) - yForLinePlot.push(currentBest) - } - currentBest = v - xForLinePlot.push(getAxisX(t)) - yForLinePlot.push(v) } - } - xForLinePlot.push(getAxisX(filteredTrials[filteredTrials.length - 1])) - yForLinePlot.push(yForLinePlot[yForLinePlot.length - 1]) - - const plotData: Partial[] = [ - { - x: filteredTrials.map(getAxisX), - y: filteredTrials.map( - (t: Optuna.Trial): number => getValue(t, objectiveId) as number + plotData.push({ + x: feasibleTrials.map(getAxisX), + y: feasibleTrials.map( + (t: Optuna.Trial): number => target.getTargetValue(t) as number ), - name: "Objective Value", + name: `${target.toLabel(h.metric_names)} of ${h.study_name}`, + marker: { + size: markerSize, + }, mode: "markers", type: "scatter", - }, - { - x: xForLinePlot, - y: yForLinePlot, - name: "Best Value", - mode: "lines", + }) + + const objectiveId = target.getObjectiveId() + if (objectiveId !== null) { + const xForLinePlot: (number | Date)[] = [] + const yForLinePlot: number[] = [] + let currentBest: number | null = null + for (let i = 0; i < feasibleTrials.length; i++) { + const t = feasibleTrials[i] + const value = target.getTargetValue(t) as number + if (t.state !== "Complete") { + continue + } + + if (currentBest === null) { + currentBest = value + xForLinePlot.push(getAxisX(t)) + yForLinePlot.push(value) + } else if ( + h.directions[objectiveId] === "maximize" && + value > currentBest + ) { + const p = feasibleTrials[i - 1] + if (!xForLinePlot.includes(getAxisX(p))) { + xForLinePlot.push(getAxisX(p)) + yForLinePlot.push(currentBest) + } + currentBest = value + xForLinePlot.push(getAxisX(t)) + yForLinePlot.push(value) + } else if ( + h.directions[objectiveId] === "minimize" && + value < currentBest + ) { + const p = feasibleTrials[i - 1] + if (!xForLinePlot.includes(getAxisX(p))) { + xForLinePlot.push(getAxisX(p)) + yForLinePlot.push(currentBest) + } + currentBest = value + xForLinePlot.push(getAxisX(t)) + yForLinePlot.push(value) + } + } + if (h.trials.length !== 0) { + xForLinePlot.push(getAxisX(h.trials[h.trials.length - 1])) + yForLinePlot.push(yForLinePlot[yForLinePlot.length - 1]) + } + plotData.push({ + x: xForLinePlot, + y: yForLinePlot, + name: `Best Value of ${h.study_name}`, + mode: "lines", + type: "scatter", + }) + } + infeasiblePlotData.push({ + x: infeasibleTrials.map(getAxisX), + y: infeasibleTrials.map( + (t: Optuna.Trial): number => target.getTargetValue(t) as number + ), + name: `Infeasible Trial of ${h.study_name}`, + marker: { + size: markerSize, + color: colorTheme === plotlyDarkTemplate ? "#666666" : "#cccccc", + }, + mode: "markers", type: "scatter", - }, - ] - plotly.react(plotDomId, plotData, layout) + showlegend: false, + }) + } + plotData.push(...infeasiblePlotData) + return plotly.react(plotDomId, plotData, layout) } diff --git a/tslib/react/test/PlotHistory.test.tsx b/tslib/react/test/PlotHistory.test.tsx index 3d121fae..15258406 100644 --- a/tslib/react/test/PlotHistory.test.tsx +++ b/tslib/react/test/PlotHistory.test.tsx @@ -18,7 +18,7 @@ describe("PlotHistory Tests", async () => { }) =>
{children}
return render( - + ) }