mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-09-11 12:30:25 +08:00
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:
@@ -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)
|
||||
}
|
||||
@@ -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>
|
||||
)
|
||||
|
||||
@@ -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
@@ -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[]) => {
|
||||
|
||||
Vendored
+10
-3
@@ -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 = {
|
||||
|
||||
Reference in New Issue
Block a user