diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index f3e405d4..9503a75a 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -6,7 +6,7 @@ from typing import TYPE_CHECKING from typing import Union import numpy as np -from optuna.distributions import BaseDistribution +from optuna.distributions import BaseDistribution, CategoricalDistribution from optuna.distributions import FloatDistribution from optuna.distributions import IntDistribution from optuna.study import StudySummary @@ -42,6 +42,41 @@ if TYPE_CHECKING: }, ) + FloatDistributionJSON = TypedDict( + "FloatDistributionJSON", + { + "type": Literal["FloatDistribution"], + "low": float, + "high": float, + "step": float, + "log": bool, + }, + ) + IntDistributionJSON = TypedDict( + "IntDistributionJSON", + { + "type": Literal["IntDistribution"], + "low": int, + "high": int, + "step": int, + "log": bool, + }, + ) + CategoricalDistributionChoiceJSON = TypedDict( + "CategoricalDistributionChoiceJSON", + { + "pytype": str, + "value": str, + } + ) + CategoricalDistributionJSON = TypedDict( + "CategoricalDistributionJSON", + { + "choices": list[CategoricalDistributionChoiceJSON] + }, + ) + DistributionJSON = Union[FloatDistributionJSON, IntDistributionJSON, CategoricalDistributionJSON] + MAX_ATTR_LENGTH = 1024 @@ -111,12 +146,22 @@ def serialize_study_detail( def serialize_frozen_trial( study_id: int, trial: FrozenTrial, study_system_attrs: dict[str, Any] ) -> dict[str, Any]: + params = [] + for param_name, param_external_value in trial.params.items(): + distribution = trial.distributions[param_name] + params.append({ + "name": param_name, + "param_internal_value": distribution.to_internal_repr(param_external_value), + "param_external_value": str(param_external_value), + "param_external_pytyp": str(type(param_external_value)), + "distribution": serialize_distribution(distribution) + }) serialized = { "trial_id": trial._trial_id, "study_id": study_id, "number": trial.number, "state": trial.state.name.capitalize(), - "params": [{"name": name, "value": str(value)} for name, value in trial.params.items()], + "params": params, "user_attrs": serialize_attrs(trial.user_attrs), "system_attrs": serialize_attrs(getattr(trial, "_system_attrs", {})), "note": note.get_note_from_system_attrs(study_system_attrs, trial._trial_id), @@ -160,6 +205,38 @@ def serialize_frozen_trial( return serialized +def serialize_distribution(distribution: BaseDistribution) -> DistributionJSON: + distribution = normalize_distribution(distribution) + if isinstance(distribution, FloatDistribution): + return { + "type": "FloatDistribution", + "low": distribution.low, + "high": distribution.high, + "step": distribution.step, + "log": distribution.log, + } + if isinstance(distribution, IntDistribution): + return { + "type": "IntDistribution", + "low": distribution.low, + "high": distribution.high, + "step": distribution.step, + "log": distribution.log, + } + if isinstance(distribution, CategoricalDistribution): + return { + "type": "CategoricalDistribution", + "choices": [ + { + "pytype": str(type(choice)), + "value": str(choice) + } + for choice in distribution.choices + ], + } + raise ValueError(f"Unexpected distribution {str(distribution)}") + + def normalize_distribution(distribution: BaseDistribution) -> BaseDistribution: if distribution.__class__.__name__ == "UniformDistribution": return FloatDistribution( @@ -200,12 +277,10 @@ def serialize_search_space( ) -> list[dict[str, Any]]: serialized = [] for param_name, distribution in search_space: - distribution = normalize_distribution(distribution) serialized.append( { "name": param_name, - "distribution": distribution.__class__.__name__, - "attributes": distribution._asdict(), + "distribution": serialize_distribution(distribution), } ) return serialized diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index f98cbd18..5b3fc2fe 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -44,8 +44,8 @@ interface StudyDetailResponse { directions: StudyDirection[] trials: TrialResponse[] best_trials: TrialResponse[] - intersection_search_space: SearchSpace[] - union_search_space: SearchSpace[] + intersection_search_space: SearchSpaceItem[] + union_search_space: SearchSpaceItem[] union_user_attrs: AttributeSpec[] has_intermediate_values: boolean note: Note diff --git a/optuna_dashboard/ts/components/BestTrialsCard.tsx b/optuna_dashboard/ts/components/BestTrialsCard.tsx index c1a5d7b4..a76802e8 100644 --- a/optuna_dashboard/ts/components/BestTrialsCard.tsx +++ b/optuna_dashboard/ts/components/BestTrialsCard.tsx @@ -41,7 +41,10 @@ export const BestTrialsCard: FC<{ Params = [ - {bestTrial.params.map((p) => `${p.name}: ${p.value}`).join(", ")}] + {bestTrial.params + .map((p) => `${p.name}: ${p.param_external_value}`) + .join(", ")} + ] Intermediate Values = [ @@ -101,7 +104,7 @@ export const BestTrialsCard: FC<{ Params = [ {trial.params - .map((p) => `${p.name}: ${p.value}`) + .map((p) => `${p.name}: ${p.param_external_value}`) .join(", ")} ] diff --git a/optuna_dashboard/ts/components/GraphContour.tsx b/optuna_dashboard/ts/components/GraphContour.tsx index b8301e20..06cb233b 100644 --- a/optuna_dashboard/ts/components/GraphContour.tsx +++ b/optuna_dashboard/ts/components/GraphContour.tsx @@ -1,5 +1,5 @@ import * as plotly from "plotly.js-dist-min" -import React, { FC, useEffect, useState } from "react" +import React, { FC, useEffect, useMemo, useState } from "react" import { Grid, FormControl, @@ -33,31 +33,44 @@ type AxisInfo = { const PADDING_RATIO = 0.05 const plotDomId = "graph-contour" +const useSearchSpace = ( + unionSearchSpaces?: SearchSpaceItem[] +): SearchSpaceItem[] => + useMemo( + () => + Array.from(unionSearchSpaces || []).sort((a, b) => + a.name > b.name ? 1 : a.name < b.name ? -1 : 0 + ), + [unionSearchSpaces] + ) + export const Contour: FC<{ study: StudyDetail | null }> = ({ study = null }) => { const theme = useTheme() const [objectiveId, setObjectiveId] = useState(0) - const [xParam, setXParam] = useState("") - const [yParam, setYParam] = useState("") - const paramNames = study?.union_search_space.map((s) => s.name) + const searchSpaces = useSearchSpace(study?.union_search_space) + const [xParam, setXParam] = useState(null) + const [yParam, setYParam] = useState(null) const objectiveNames: string[] = study?.objective_names || [] - if (!xParam && paramNames && paramNames.length > 0) { - setXParam(paramNames[0]) + if (xParam === null && searchSpaces.length > 0) { + setXParam(searchSpaces[0]) } - if (!yParam && paramNames && paramNames.length > 1) { - setYParam(paramNames[1]) + if (yParam === null && searchSpaces.length > 1) { + setYParam(searchSpaces[1]) } const handleObjectiveChange = (event: SelectChangeEvent) => { setObjectiveId(event.target.value as number) } const handleXParamChange = (event: SelectChangeEvent) => { - setXParam(event.target.value as string) + const param = searchSpaces.find((s) => s.name === event.target.value) + setXParam(param || null) } const handleYParamChange = (event: SelectChangeEvent) => { - setYParam(event.target.value as string) + const param = searchSpaces.find((s) => s.name === event.target.value) + setYParam(param || null) } useEffect(() => { @@ -66,7 +79,7 @@ export const Contour: FC<{ } }, [study, objectiveId, xParam, yParam, theme.palette.mode]) - const space: SearchSpace[] = study ? study.union_search_space : [] + const space: SearchSpaceItem[] = study ? study.union_search_space : [] return ( @@ -98,7 +111,7 @@ export const Contour: FC<{ x: - {space.map((d, i) => ( {d.name} @@ -108,7 +121,7 @@ export const Contour: FC<{ y: - {space.map((d, i) => ( {d.name} @@ -126,73 +139,6 @@ export const Contour: FC<{ ) } -const isNumerical = (trials: Trial[], paramName: string): boolean => { - return trials.every((t) => { - const param = t.params.find((param) => param.name === paramName) - if (!param) return true - const val = param.value - return typeof (Number(val) || val) === "number" - }) -} - -const getAxisInfo = (trials: Trial[], paramName: string): AxisInfo => { - const values = trials.map((trial) => { - const param = trial.params.find((p) => p.name === paramName) - return param ? Number(param.value) || param.value : null - }) - - let min: number - let max: number - let isLog: boolean - let isCat: boolean - - if (isNumerical(trials, paramName)) { - const minValue = Math.min(...(values as number[])) - const maxValue = Math.max(...(values as number[])) - const padding = (maxValue - minValue) * PADDING_RATIO - min = minValue - padding - max = maxValue + padding - isLog = false - isCat = false - } else { - const uniqueValues = unique(values) - const span = uniqueValues.length - (uniqueValues.includes(null) ? 2 : 1) - const padding = span * PADDING_RATIO - min = -padding - max = span + padding - isLog = false - isCat = true - } - - const indices = isNumerical(trials, paramName) - ? unique((values as (number | null)[]).filter((v) => v !== null)).sort( - (a, b) => a - b - ) - : unique((values as (string | null)[]).filter((v) => v !== null)).sort( - (a, b) => - a.toString().toLowerCase() < b.toString().toLowerCase() - ? -1 - : a.toString().toLowerCase() > b.toString().toLowerCase() - ? 1 - : 0 - ) - - if (indices.length >= 2 && isNumerical(trials, paramName)) { - indices.unshift(min) - indices.push(max) - } - - return { - name: paramName, - min, - max, - isLog, - isCat, - indices, - values, - } -} - const filterFunc = (trial: Trial, objectiveId: number): boolean => { return ( trial.state === "Complete" && @@ -205,8 +151,8 @@ const filterFunc = (trial: Trial, objectiveId: number): boolean => { const plotContour = ( study: StudyDetail, objectiveId: number, - xParam: string, - yParam: string, + xParam: SearchSpaceItem | null, + yParam: SearchSpaceItem | null, mode: string ) => { if (document.getElementById(plotDomId) === null) { @@ -215,16 +161,15 @@ const plotContour = ( const trials: Trial[] = study ? study.trials : [] const filteredTrials = trials.filter((t) => filterFunc(t, objectiveId)) - - if (filteredTrials.length === 0) { + if (filteredTrials.length === 0 || xParam === null || yParam === null) { plotly.react(plotDomId, [], { template: mode === "dark" ? plotlyDarkTemplate : {}, }) return } - const xAxis = getAxisInfo(trials, xParam) - const yAxis = getAxisInfo(trials, yParam) + const xAxis = getAxisInfo(study, trials, xParam) + const yAxis = getAxisInfo(study, trials, yParam) const xIndices = xAxis.indices const yIndices = yAxis.indices @@ -279,11 +224,11 @@ const plotContour = ( const layout: Partial = { xaxis: { - title: xParam, + title: xParam.name, type: xAxis.isCat ? "category" : undefined, }, yaxis: { - title: yParam, + title: yParam.name, type: yAxis.isCat ? "category" : undefined, }, margin: { @@ -296,3 +241,87 @@ const plotContour = ( } plotly.react(plotDomId, plotData, layout) } + +const getAxisInfoForNumericalParams = ( + trials: Trial[], + paramName: string, + distribution: FloatDistribution | IntDistribution +): AxisInfo => { + const padding = (distribution.high - distribution.low) * PADDING_RATIO + const min = distribution.low - padding + const max = distribution.high + padding + + const values = trials.map( + (trial) => + trial.params.find((p) => p.name === paramName)?.param_internal_value || + null + ) + const indices = unique(values) + .filter((v) => v !== null) + .sort((a, b) => a - b) + if (indices.length >= 2) { + indices.unshift(min) + indices.push(max) + } + return { + name: paramName, + min, + max, + isLog: distribution.log, + isCat: false, + indices, + values, + } +} + +const getAxisInfoForCategoricalParams = ( + trials: Trial[], + paramName: string, + distribution: CategoricalDistribution +): AxisInfo => { + const values = trials.map( + (trial) => + trial.params.find((p) => p.name === paramName)?.param_external_value || + null + ) + const isDynamic = values.some((v) => v === null) + const span = distribution.choices.length - (isDynamic ? 2 : 1) + const padding = span * PADDING_RATIO + const min = -padding + const max = span + padding + + const indices = distribution.choices + .map((c) => c.value) + .sort((a, b) => + a.toLowerCase() < b.toLowerCase() + ? -1 + : a.toLowerCase() > b.toLowerCase() + ? 1 + : 0 + ) + return { + name: paramName, + min, + max, + isLog: false, + isCat: true, + indices, + values, + } +} + +const getAxisInfo = ( + study: StudyDetail, + trials: Trial[], + param: SearchSpaceItem +): AxisInfo => { + if (param.distribution.type === "CategoricalDistribution") { + return getAxisInfoForCategoricalParams( + trials, + param.name, + param.distribution + ) + } else { + return getAxisInfoForNumericalParams(trials, param.name, param.distribution) + } +} diff --git a/optuna_dashboard/ts/components/GraphHistory.tsx b/optuna_dashboard/ts/components/GraphHistory.tsx index 4f55d735..8c38fc48 100644 --- a/optuna_dashboard/ts/components/GraphHistory.tsx +++ b/optuna_dashboard/ts/components/GraphHistory.tsx @@ -16,7 +16,11 @@ import { useTheme, } from "@mui/material" import { plotlyDarkTemplate } from "./PlotlyDarkMode" -import { useFilteredTrials, Target, useTargetList } from "../trialFilter" +import { + useFilteredTrials, + Target, + useObjectiveAndSystemAttrTargets, +} from "../trialFilter" const plotDomId = "graph-history" @@ -30,11 +34,11 @@ export const GraphHistory: FC<{ const [filterPrunedTrial, setFilterPrunedTrial] = useState(false) const objectiveNames: string[] = study?.objective_names || [] - const targetList = useTargetList(study) + const targetList = useObjectiveAndSystemAttrTargets(study) const [targetIndex, setTargetIndex] = useState(0) const trials = useFilteredTrials( study, - targetList[targetIndex], + [targetList[targetIndex]], filterCompleteTrial, filterPrunedTrial ) diff --git a/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx b/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx index ca394e62..7347f42e 100644 --- a/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx +++ b/optuna_dashboard/ts/components/GraphParallelCoordinate.tsx @@ -136,15 +136,10 @@ const plotCoordinate = ( }, ] study.intersection_search_space.forEach((s) => { - const valueStrings = filteredTrials.map((t) => { - const param = t.params.find((p) => p.name === s.name) - return param!.value - }) - const isnum = valueStrings.every((v) => { - return !isNaN(Number(v)) - }) - if (isnum) { - const values: number[] = valueStrings.map((v) => parseFloat(v)) + const values: number[] = filteredTrials.map( + (t) => t.params.find((p) => p.name === s.name)!.param_internal_value + ) + if (s.distribution.type !== "CategoricalDistribution") { dimensions.push({ label: breakLabelIfTooLong(s.name), values: values, @@ -152,11 +147,7 @@ const plotCoordinate = ( }) } else { // categorical - const vocabSet = new Set(valueStrings) - const vocabArr = Array.from(vocabSet) - const values: number[] = valueStrings.map((v) => - vocabArr.findIndex((vocab) => v === vocab) - ) + const vocabArr: string[] = s.distribution.choices.map((c) => c.value) const tickvals: number[] = vocabArr.map((v, i) => i) dimensions.push({ label: breakLabelIfTooLong(s.name), diff --git a/optuna_dashboard/ts/components/GraphParetoFront.tsx b/optuna_dashboard/ts/components/GraphParetoFront.tsx index ef3dd646..22c4007c 100644 --- a/optuna_dashboard/ts/components/GraphParetoFront.tsx +++ b/optuna_dashboard/ts/components/GraphParetoFront.tsx @@ -105,7 +105,7 @@ const makeHovertext = (trial: Trial): string => { number: trial.number, values: trial.values, params: trial.params - .map((p) => [p.name, p.value]) + .map((p) => [p.name, p.param_external_value]) .reduce((obj, [key, value]) => ({ ...obj, [key]: value }), {}), }, undefined, diff --git a/optuna_dashboard/ts/components/GraphSlice.tsx b/optuna_dashboard/ts/components/GraphSlice.tsx index 31231881..d1e201bf 100644 --- a/optuna_dashboard/ts/components/GraphSlice.tsx +++ b/optuna_dashboard/ts/components/GraphSlice.tsx @@ -1,5 +1,5 @@ import * as plotly from "plotly.js-dist-min" -import React, { ChangeEvent, FC, useEffect, useState } from "react" +import React, { ChangeEvent, FC, useEffect, useMemo, useState } from "react" import { Grid, FormControl, @@ -13,52 +13,76 @@ import { Box, } from "@mui/material" import { plotlyDarkTemplate } from "./PlotlyDarkMode" +import { Target, useFilteredTrials, useObjectiveTargets } from "../trialFilter" +import { useSnackbar } from "notistack" const plotDomId = "graph-slice" -// TODO(c-bata): Check `log` field of IntDistribution and FloatDistribution. -const logDistributions = ["LogUniformDistribution", "IntLogUniformDistribution"] +const useSearchSpace = ( + unionSearchSpaces?: SearchSpaceItem[] +): SearchSpaceItem[] => + useMemo( + () => + Array.from(unionSearchSpaces || []).sort((a, b) => + a.name > b.name ? 1 : a.name < b.name ? -1 : 0 + ), + [unionSearchSpaces] + ) + +const isLogScale = (s: SearchSpaceItem): boolean => { + if (s.distribution.type === "CategoricalDistribution") { + return false + } + return s.distribution.log +} export const GraphSlice: FC<{ study: StudyDetail | null }> = ({ study = null }) => { const theme = useTheme() - const trials: Trial[] = study !== null ? study.trials : [] + const { enqueueSnackbar } = useSnackbar() + const [objectiveId, setObjectiveId] = useState(0) - const [selected, setSelected] = useState(null) - const [logXScale, setLogXScale] = useState(false) + const [selected, setSelected] = useState(null) const [logYScale, setLogYScale] = useState(false) - const paramNames = study?.union_search_space.map((s) => s.name) - const distributions = new Map( - study?.union_search_space.map((s) => [s.name, s.distribution]) - ) + const searchSpaces = useSearchSpace(study?.union_search_space) + + const targets = useObjectiveTargets(study) + const filterTargets: Target[] = [targets[objectiveId]] + if (selected !== null) filterTargets.push(new Target("params", selected.name)) + const trials = useFilteredTrials(study, filterTargets, false, false) + const objectiveNames: string[] = study?.objective_names || [] - if (selected === null && paramNames && paramNames.length > 0) { - const distribution = distributions.get(paramNames[0]) || "" - setSelected(paramNames[0]) - setLogXScale(logDistributions.includes(distribution)) + if (selected === null && searchSpaces.length > 0) { + setSelected(searchSpaces[0]) } useEffect(() => { plotSlice( trials, - objectiveId, + targets[objectiveId], selected, - logXScale, logYScale, theme.palette.mode ) - }, [trials, objectiveId, selected, logXScale, logYScale, theme.palette.mode]) + }, [trials, targets[objectiveId], selected, logYScale, theme.palette.mode]) const handleObjectiveChange = (event: SelectChangeEvent) => { setObjectiveId(event.target.value as number) } const handleSelectedParam = (e: SelectChangeEvent) => { - const paramName = e.target.value - const distribution = distributions.get(paramName) || "" - setSelected(paramName) - setLogXScale(logDistributions.includes(distribution)) + const s = searchSpaces.find((s) => s.name === e.target.value) + if (s === undefined) { + enqueueSnackbar( + `Cannot find ${e.target.value} param in the search space.`, + { + variant: "error", + } + ) + return + } + setSelected(s) } const handleLogYScaleChange = (e: ChangeEvent) => { @@ -93,9 +117,9 @@ export const GraphSlice: FC<{ )} Parameter: - + {searchSpaces?.map((s, i) => ( + {objectiveNames.length === study?.directions.length ? objectiveNames[i] : `${i}`} @@ -119,32 +143,10 @@ export const GraphSlice: FC<{ ) } -const filterFunc = ( - trial: Trial, - objectiveId: number, - selected: string | null -): boolean => { - if (trial.state !== "Complete" && trial.state !== "Pruned") { - return false - } - if (trial.params.find((p) => p.name == selected) === undefined) { - return false - } - if (trial.values === undefined) { - return false - } - return ( - trial.values.length > objectiveId && - trial.values[objectiveId] !== "inf" && - trial.values[objectiveId] !== "-inf" - ) -} - const plotSlice = ( trials: Trial[], - objectiveId: number, - selected: string | null, - logXScale: boolean, + target: Target, + selected: SearchSpaceItem | null, logYScale: boolean, mode: string ) => { @@ -160,8 +162,8 @@ const plotSlice = ( b: 0, }, xaxis: { - title: selected || "", - type: logXScale ? "log" : "linear", + title: selected?.name || "", + type: selected !== null && isLogScale(selected) ? "log" : "linear", gridwidth: 1, automargin: true, }, @@ -174,34 +176,27 @@ const plotSlice = ( showlegend: false, template: mode === "dark" ? plotlyDarkTemplate : {}, } - - const filteredTrials = trials.filter((t) => - filterFunc(t, objectiveId, selected) - ) - - if (filteredTrials.length === 0 || selected === null) { + if (selected === null) { + plotly.react(plotDomId, [], layout) + return + } + if (trials.length === 0) { plotly.react(plotDomId, [], layout) return } - const objectiveValues: number[] = filteredTrials.map( - (t) => t.values![objectiveId] as number + const objectiveValues: number[] = trials.map( + (t) => target.getTargetValue(t) as number ) - const valueStrings = filteredTrials.map((t) => { - return t.params.find((p) => p.name == selected)!.value - }) + const paramTarget = new Target("params", selected.name) + const values = trials.map((t) => paramTarget.getTargetValue(t) as number) - const trialNumbers: number[] = filteredTrials.map((t) => t.number) - - const isnum = valueStrings.every((v) => { - return !isNaN(Number(v)) - }) - if (isnum) { - const valuesNum: number[] = valueStrings.map((v) => parseFloat(v)) + const trialNumbers: number[] = trials.map((t) => t.number) + if (selected.distribution.type !== "CategoricalDistribution") { const trace: plotly.Data[] = [ { type: "scatter", - x: valuesNum, + x: values, y: objectiveValues, mode: "markers", marker: { @@ -219,23 +214,19 @@ const plotSlice = ( }, ] layout["xaxis"] = { - title: selected, - type: logXScale ? "log" : "linear", + title: selected.name, + type: selected.distribution.log ? "log" : "linear", gridwidth: 1, automargin: true, // Otherwise the label is outside of the plot } plotly.react(plotDomId, trace, layout) } else { - const vocabSet = new Set(valueStrings) - const vocabArr = Array.from(vocabSet) - const valuesCategorical: number[] = valueStrings.map((v) => - vocabArr.findIndex((vocab) => v === vocab) - ) + const vocabArr = selected.distribution.choices.map((c) => c.value) const tickvals: number[] = vocabArr.map((v, i) => i) const trace: plotly.Data[] = [ { type: "scatter", - x: valuesCategorical, + x: values, y: objectiveValues, mode: "markers", marker: { @@ -253,8 +244,8 @@ const plotSlice = ( }, ] layout["xaxis"] = { - title: selected, - type: logXScale ? "log" : "linear", + title: selected.name, + type: "linear", gridwidth: 1, tickvals: tickvals, ticktext: vocabArr, diff --git a/optuna_dashboard/ts/components/TrialList.tsx b/optuna_dashboard/ts/components/TrialList.tsx index 44b180c7..8e20df96 100644 --- a/optuna_dashboard/ts/components/TrialList.tsx +++ b/optuna_dashboard/ts/components/TrialList.tsx @@ -162,7 +162,10 @@ const TrialListDetail: FC<{ Params = [ - {trial.params.map((p) => `${p.name}: ${p.value}`).join(", ")}] + {trial.params + .map((p) => `${p.name}: ${p.param_external_value}`) + .join(", ")} + ] Started At ={" "} diff --git a/optuna_dashboard/ts/components/TrialTable.tsx b/optuna_dashboard/ts/components/TrialTable.tsx index ba271bf1..f3a55922 100644 --- a/optuna_dashboard/ts/components/TrialTable.tsx +++ b/optuna_dashboard/ts/components/TrialTable.tsx @@ -139,20 +139,23 @@ export const TrialTable: FC<{ studyDetail?.intersection_search_space.length ) { studyDetail?.intersection_search_space.forEach((s) => { - const sortable = s.distribution !== "CategoricalDistribution" - const filterable = s.distribution === "CategoricalDistribution" + const sortable = s.distribution.type !== "CategoricalDistribution" + const filterable = s.distribution.type === "CategoricalDistribution" columns.push({ field: "params", label: `Param ${s.name}`, toCellValue: (i) => - trials[i].params.find((p) => p.name === s.name)?.value || null, + trials[i].params.find((p) => p.name === s.name) + ?.param_external_value || null, sortable: sortable, filterable: filterable, less: (firstEl, secondEl): number => { - const firstVal = firstEl.params.find((p) => p.name === s.name)?.value + const firstVal = firstEl.params.find( + (p) => p.name === s.name + )?.param_internal_value const secondVal = secondEl.params.find( (p) => p.name === s.name - )?.value + )?.param_internal_value if (firstVal === secondVal) { return 0 @@ -171,7 +174,9 @@ export const TrialTable: FC<{ field: "params", label: "Params", toCellValue: (i) => - trials[i].params.map((p) => p.name + ": " + p.value).join(", "), + trials[i].params + .map((p) => p.name + ": " + p.param_external_value) + .join(", "), }) } diff --git a/optuna_dashboard/ts/trialFilter.ts b/optuna_dashboard/ts/trialFilter.ts index 9952411c..1e249761 100644 --- a/optuna_dashboard/ts/trialFilter.ts +++ b/optuna_dashboard/ts/trialFilter.ts @@ -1,10 +1,12 @@ import { useMemo } from "react" +type TargetKind = "objective" | "user_attr" | "params" + export class Target { - kind: "objective" | "user_attr" + kind: TargetKind key: number | string - constructor(kind: "objective" | "user_attr", key: number | string) { + constructor(kind: TargetKind, key: number | string) { this.kind = kind this.key = key } @@ -18,8 +20,10 @@ export class Target { if (typeof this.key !== "string") { return false } - } else { - return false + } else if (this.kind === "params") { + if (typeof this.key !== "string") { + return false + } } return true } @@ -31,12 +35,17 @@ export class Target { return objectiveNames[objectiveId] } return `Objective ${objectiveId}` - } else { + } else if (this.kind === "user_attr") { return `User Attribute ${this.key}` + } else { + return `Param ${this.key}` } } getObjectiveId(): number | null { + if (this.kind !== "objective") { + return null + } return this.key as number } @@ -68,6 +77,12 @@ export class Target { return null } return value + } else if (this.kind === "params") { + const param = trial.params.find((p) => p.name === this.key) + if (param === undefined) { + return null + } + return param.param_internal_value } return null } @@ -75,7 +90,7 @@ export class Target { export const useFilteredTrials = ( study: StudyDetail | null, - target: Target, + targets: Target[], filterComplete: boolean, filterPruned: boolean ): Trial[] => @@ -93,11 +108,22 @@ export const useFilteredTrials = ( if (t.state === "Pruned" && filterPruned) { return false } - return target.getTargetValue(t) !== null + return targets.every((target) => target.getTargetValue(t) !== null) }) - }, [study?.trials, target, filterComplete, filterPruned]) + }, [study?.trials, targets, filterComplete, filterPruned]) -export const useTargetList = (study: StudyDetail | null): Target[] => +export const useObjectiveTargets = (study: StudyDetail | null): Target[] => + useMemo(() => { + if (study !== null) { + return study.directions.map((v, i) => new Target("objective", i)) + } else { + return [new Target("objective", 0)] + } + }, [study?.directions]) + +export const useObjectiveAndSystemAttrTargets = ( + study: StudyDetail | null +): Target[] => useMemo(() => { if (study !== null) { return [ diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index 425045a2..440789ba 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -11,10 +11,32 @@ type TrialValueNumber = number | "inf" | "-inf" type TrialIntermediateValueNumber = number | "inf" | "-inf" | "nan" type TrialState = "Running" | "Complete" | "Pruned" | "Fail" | "Waiting" type StudyDirection = "maximize" | "minimize" | "not_set" + +type FloatDistribution = { + type: "FloatDistribution" + low: number + high: number + step: number + log: boolean +} + +type IntDistribution = { + type: "IntDistribution" + low: number + high: number + step: number + log: boolean +} + +type CategoricalDistribution = { + type: "CategoricalDistribution" + choices: { pytype: string; value: string }[] +} + type Distribution = - | "FloatDistribution" - | "IntDistribution" - | "CategoricalDistribution" + | FloatDistribution + | IntDistribution + | CategoricalDistribution type GraphVisibility = { history: boolean @@ -34,7 +56,10 @@ type TrialIntermediateValue = { type TrialParam = { name: string - value: string + param_internal_value: number + param_external_value: string + param_external_type: string + distribution: Distribution } type ParamImportance = { @@ -43,7 +68,7 @@ type ParamImportance = { distribution: Distribution } -type SearchSpace = { +type SearchSpaceItem = { name: string distribution: Distribution } @@ -94,8 +119,8 @@ type StudyDetail = { datetime_start: Date best_trials: Trial[] trials: Trial[] - intersection_search_space: SearchSpace[] - union_search_space: SearchSpace[] + intersection_search_space: SearchSpaceItem[] + union_search_space: SearchSpaceItem[] union_user_attrs: AttributeSpec[] has_intermediate_values: boolean note: Note diff --git a/typescript_tests/TrialTable.test.tsx b/typescript_tests/TrialTable.test.tsx index c443df13..df772179 100644 --- a/typescript_tests/TrialTable.test.tsx +++ b/typescript_tests/TrialTable.test.tsx @@ -61,21 +61,21 @@ const studyDetail: StudyDetail = { intersection_search_space: [ { name: "x", - distribution: "FloatDistribution" as Distribution, + distribution: "FloatDistribution" as DistributionName, }, { name: "y", - distribution: "FloatDistribution" as Distribution, + distribution: "FloatDistribution" as DistributionName, }, ], union_search_space: [ { name: "x", - distribution: "FloatDistribution" as Distribution, + distribution: "FloatDistribution" as DistributionName, }, { name: "y", - distribution: "FloatDistribution" as Distribution, + distribution: "FloatDistribution" as DistributionName, }, ], union_user_attrs: [