Merge pull request #575 from c-bata/standalone-app-union-user-attrs

Support external_param_value and union_user_attrs on optuna-dashboard-wasm.
This commit is contained in:
c-bata
2023-08-27 03:14:06 +09:00
committed by GitHub
5 changed files with 308 additions and 63 deletions
@@ -0,0 +1,100 @@
import * as plotly from "plotly.js-dist-min"
import React, { FC, useEffect } from "react"
import { Box, Typography, useTheme, CardContent, Card } from "@mui/material"
import { plotlyDarkTemplate } from "../PlotlyDarkMode"
const plotDomId = "graph-intermediate-values"
export const PlotIntermediateValues: FC<{
trials: Trial[]
includePruned: boolean
logScale: boolean
}> = ({ trials, includePruned, logScale }) => {
const theme = useTheme()
useEffect(() => {
plotIntermediateValue(
trials,
theme.palette.mode,
false,
!includePruned,
logScale
)
}, [trials, theme.palette.mode, false, includePruned, logScale])
return (
<Card>
<CardContent>
<Typography
variant="h6"
sx={{ margin: "1em 0", fontWeight: theme.typography.fontWeightBold }}
>
Intermediate values
</Typography>
<Box id={plotDomId} sx={{ height: "450px" }} />
</CardContent>
</Card>
)
}
const plotIntermediateValue = (
trials: Trial[],
mode: string,
filterCompleteTrial: boolean,
filterPrunedTrial: boolean,
logScale: boolean
) => {
if (document.getElementById(plotDomId) === null) {
return
}
const layout: Partial<plotly.Layout> = {
margin: {
l: 50,
t: 0,
r: 50,
b: 0,
},
yaxis: {
title: "Objective Value",
type: logScale ? "log" : "linear",
},
xaxis: {
title: "Step",
type: "linear",
},
uirevision: "true",
template: mode === "dark" ? plotlyDarkTemplate : {},
}
if (trials.length === 0) {
plotly.react(plotDomId, [], layout)
return
}
const filteredTrials = trials.filter(
(t) =>
(!filterCompleteTrial && t.state === "Complete") ||
(!filterPrunedTrial &&
t.state === "Pruned" &&
t.values &&
t.values.length > 0) ||
t.state == "Running"
)
const plotData: Partial<plotly.PlotData>[] = filteredTrials.map((trial) => {
const values = trial.intermediate_values.filter(
(iv) => iv.value !== "inf" && iv.value !== "-inf" && iv.value !== "nan"
)
return {
x: values.map((iv) => iv.step),
y: values.map((iv) => iv.value),
marker: { maxdisplayed: 10 },
mode: "lines+markers",
type: "scatter",
name:
trial.state !== "Running"
? `trial #${trial.number}`
: `trial #${trial.number} (running)`,
}
})
plotly.react(plotDomId, plotData, layout)
}
+26 -7
View File
@@ -11,6 +11,7 @@ import {
Card,
CardContent,
} from "@mui/material"
import Grid2 from "@mui/material/Unstable_Grid2"
import { Home } from "@mui/icons-material"
import Brightness4Icon from "@mui/icons-material/Brightness4"
import Brightness7Icon from "@mui/icons-material/Brightness7"
@@ -19,6 +20,7 @@ import { studiesState } from "../state"
import { TrialTable } from "./TrialTable"
import { PlotHistory } from "./PlotHistory"
import { PlotImportance } from "./PlotImportance"
import { PlotIntermediateValues } from "./PlotIntermediateValues"
const useStudyValue = (idx: number): Study | null => {
const studies = useRecoilValue<Study[]>(studiesState)
@@ -83,7 +85,7 @@ export const StudyDetail: FC<{
},
}}
>
<div>
<>
<Typography
variant="h4"
sx={{
@@ -102,17 +104,34 @@ export const StudyDetail: FC<{
<PlotHistory study={study} />
</CardContent>
</Card>
<Card sx={{ margin: theme.spacing(2) }}>
<CardContent>
{!!study && <PlotImportance study={study} />}
</CardContent>
</Card>
<Grid2 container spacing={0}>
<Grid2 xs={6}>
<Card sx={{ margin: theme.spacing(2) }}>
<CardContent>
{!!study && <PlotImportance study={study} />}
</CardContent>
</Card>
</Grid2>
<Grid2 xs={6}>
<Card sx={{ margin: theme.spacing(2) }}>
<CardContent>
{!!study && (
<PlotIntermediateValues
trials={study.trials}
includePruned={false}
logScale={false}
/>
)}
</CardContent>
</Card>
</Grid2>
</Grid2>
<Card sx={{ margin: theme.spacing(2) }}>
<CardContent>
{!!study && <TrialTable study={study} initialRowsPerPage={10} />}
</CardContent>
</Card>
</div>
</>
</Container>
</div>
)
+60 -30
View File
@@ -20,36 +20,6 @@ export const TrialTable: FC<{
},
]
study.union_search_space.forEach((s) => {
columns.push({
field: "params",
label: `Param ${s.name}`,
toCellValue: (i) =>
trials[i].params.find((p) => p.name === s.name)?.param_internal_value ||
null,
sortable: true,
filterable: false,
less: (firstEl, secondEl): number => {
const firstVal = firstEl.params.find(
(p) => p.name === s.name
)?.param_internal_value
const secondVal = secondEl.params.find(
(p) => p.name === s.name
)?.param_internal_value
if (firstVal === secondVal) {
return 0
} else if (firstVal && secondVal) {
return firstVal < secondVal ? 1 : -1
} else if (firstVal) {
return -1
} else {
return 1
}
},
})
})
if (study === null || study.directions.length == 1) {
columns.push({
field: "values",
@@ -117,6 +87,66 @@ export const TrialTable: FC<{
columns.push(...objectiveColumns)
}
study.union_search_space.forEach((s) => {
columns.push({
field: "params",
label: `Param ${s.name}`,
toCellValue: (i) =>
trials[i].params.find((p) => p.name === s.name)?.param_external_value ??
null,
sortable: true,
filterable: false,
less: (firstEl, secondEl): number => {
const firstVal = firstEl.params.find(
(p) => p.name === s.name
)?.param_internal_value
const secondVal = secondEl.params.find(
(p) => p.name === s.name
)?.param_internal_value
if (firstVal === secondVal) {
return 0
} else if (firstVal && secondVal) {
return firstVal < secondVal ? 1 : -1
} else if (firstVal) {
return -1
} else {
return 1
}
},
})
})
study.union_user_attrs.forEach((attr_spec) => {
columns.push({
field: "user_attrs",
label: `UserAttribute ${attr_spec.key}`,
toCellValue: (i) =>
trials[i].user_attrs.find((attr) => attr.key === attr_spec.key)
?.value || null,
sortable: attr_spec.sortable,
filterable: false,
less: (firstEl, secondEl): number => {
const firstVal = firstEl.user_attrs.find(
(attr) => attr.key === attr_spec.key
)?.value
const secondVal = secondEl.user_attrs.find(
(attr) => attr.key === attr_spec.key
)?.value
if (firstVal === secondVal) {
return 0
} else if (firstVal && secondVal) {
return firstVal < secondVal ? 1 : -1
} else if (firstVal) {
return -1
} else {
return 1
}
},
})
})
return (
<DataGrid<Trial>
columns={columns}
+112 -23
View File
@@ -80,30 +80,25 @@ const getStudies = (db: SQLite3DB): Study[] => {
const trials = getTrials(db, studyId)
const union_search_space: SearchSpaceItem[] = []
const union_user_attrs: AttributeSpec[] = []
let intersection_search_space: Set<SearchSpaceItem> = new Set()
trials.forEach((trial) => {
const params: TrialParam[] = []
const param_names = new Set<string>()
db.exec({
sql:
"SELECT param_name, param_value" +
` FROM trial_params WHERE trial_id = ${trial.trial_id}`,
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const param_name = vals[0]
// TODO(c-bata): Support param_external_value
params.push({
name: param_name,
param_internal_value: vals[1],
})
const userAttrs = getTrialUserAttributes(db, trial.trial_id)
userAttrs.forEach((attr) => {
if (union_user_attrs.findIndex((s) => s.key === attr.key) == -1) {
union_user_attrs.push({ key: attr.key, sortable: false })
}
})
param_names.add(param_name)
if (
union_search_space.findIndex((s) => s.name === param_name) == -1
) {
union_search_space.push({ name: param_name })
}
},
const params = getTrialParams(db, trial.trial_id)
const param_names = new Set<string>()
params.forEach((param) => {
param_names.add(param.name)
if (
union_search_space.findIndex((s) => s.name === param.name) == -1
) {
union_search_space.push({ name: param.name })
}
})
if (intersection_search_space.size === 0) {
param_names.forEach((s) => {
@@ -119,6 +114,7 @@ const getStudies = (db: SQLite3DB): Study[] => {
)
}
trial.params = params
trial.user_attrs = userAttrs
})
if (objective === 0) {
@@ -128,6 +124,7 @@ const getStudies = (db: SQLite3DB): Study[] => {
directions: [direction],
union_search_space: union_search_space,
intersection_search_space: Array.from(intersection_search_space),
union_user_attrs: union_user_attrs,
trials: trials,
})
return
@@ -164,8 +161,9 @@ const getTrials = (db: SQLite3DB, studyId: number): Trial[] => {
study_id: studyId,
state: state,
values: getTrialValues(db, trialId),
params: [], // Set this column later
intermediate_values: getTrialIntermediateValues(db, trialId),
params: [], // Set this column later
user_attrs: [], // Set this column later
datetime_start: vals[3],
datetime_complete: vals[4],
}
@@ -196,6 +194,97 @@ const getTrialValues = (db: SQLite3DB, trialId: number): TrialValueNumber[] => {
return values
}
const getTrialParams = (db: SQLite3DB, trialId: number): TrialParam[] => {
const params: TrialParam[] = []
db.exec({
sql:
"SELECT param_name, param_value, distribution_json" +
` FROM trial_params WHERE trial_id = ${trialId}`,
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const distribution = parseDistributionJSON(vals[2])
params.push({
name: vals[0],
param_internal_value: vals[1],
param_external_type: distribution.type,
param_external_value: paramInternalValueToExternalValue(
distribution,
vals[1]
),
distribution: distribution,
})
},
})
return params
}
const paramInternalValueToExternalValue = (
distribution: Distribution,
internalValue: number
): string => {
if (distribution.type === "FloatDistribution") {
return internalValue.toString()
} else if (distribution.type === "IntDistribution") {
return internalValue.toString()
} else {
return distribution.choices[internalValue].value
}
}
const parseDistributionJSON = (t: string): Distribution => {
const parsed = JSON.parse(t)
if (parsed.name === "FloatDistribution") {
return {
type: "FloatDistribution",
low: parsed.attributes.low as number,
high: parsed.attributes.high as number,
step: parsed.attributes.step as number,
log: parsed.attributes.log as boolean,
}
} else if (parsed.name === "IntDistribution") {
return {
type: "IntDistribution",
low: parsed.attributes.low as number,
high: parsed.attributes.high as number,
step: parsed.attributes.step as number,
log: parsed.attributes.log as boolean,
}
} else {
// eslint-disable-next-line @typescript-eslint/no-explicit-any
const choices = parsed.attributes.choices.map((value: any) => {
// TODO(c-bata): Support other types
return {
pytype: "str",
value: value.toString(),
}
})
return {
type: "CategoricalDistribution",
choices: choices,
}
}
}
const getTrialUserAttributes = (
db: SQLite3DB,
trialId: number
): Attribute[] => {
const attrs: Attribute[] = []
db.exec({
sql:
"SELECT key, value_json" +
` FROM trial_user_attributes WHERE trial_id = ${trialId}`,
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
attrs.push({
key: vals[0],
value: vals[1],
})
},
})
return attrs
}
const getTrialIntermediateValues = (
db: SQLite3DB,
trialId: number
@@ -204,7 +293,7 @@ const getTrialIntermediateValues = (
db.exec({
sql:
"SELECT step, intermediate_value, intermediate_value_type" +
` FROM trial_values WHERE trial_id = ${trialId}` +
` FROM trial_intermediate_values WHERE trial_id = ${trialId}` +
" ORDER BY step",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
+10 -3
View File
@@ -42,12 +42,18 @@ type Attribute = {
value: string
}
type AttributeSpec = {
key: string
sortable: boolean
}
type Study = {
study_id: number
study_name: string
directions: StudyDirection[]
union_search_space: SearchSpaceItem[]
intersection_search_space: SearchSpaceItem[]
union_user_attrs: AttributeSpec[]
datetime_start?: Date
trials: Trial[]
}
@@ -60,6 +66,7 @@ type Trial = {
values?: TrialValueNumber[]
params: TrialParam[]
intermediate_values: TrialIntermediateValue[]
user_attrs: Attribute[]
datetime_start?: Date
datetime_complete?: Date
}
@@ -67,9 +74,9 @@ type Trial = {
type TrialParam = {
name: string
param_internal_value: number
// param_external_value: string
// param_external_type: string
// distribution: Distribution
param_external_value: string
param_external_type: string
distribution: Distribution
}
type SearchSpaceItem = {