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:
-
y:
-
+
{space.map((d, i) => (
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: [