diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py
index 6442fa91..7c50b6f2 100644
--- a/optuna_dashboard/_serializer.py
+++ b/optuna_dashboard/_serializer.py
@@ -7,6 +7,7 @@ from typing import Union
import numpy as np
from optuna.distributions import BaseDistribution
+from optuna.distributions import CategoricalDistribution
from optuna.study import StudySummary
from optuna.trial import FrozenTrial
@@ -40,6 +41,44 @@ 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",
+ {
+ "type": Literal["CategoricalDistribution"],
+ "choices": list[CategoricalDistributionChoiceJSON],
+ },
+ )
+ DistributionJSON = Union[
+ FloatDistributionJSON, IntDistributionJSON, CategoricalDistributionJSON
+ ]
+
MAX_ATTR_LENGTH = 1024
@@ -109,12 +148,26 @@ 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.get(param_name)
+ if distribution is None:
+ continue
+ 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),
@@ -158,6 +211,89 @@ def serialize_frozen_trial(
return serialized
+def serialize_distribution(distribution: BaseDistribution) -> DistributionJSON:
+ if distribution.__class__.__name__ == "FloatDistribution":
+ # Added from Optuna v3.0
+ float_distribution: FloatDistributionJSON = {
+ "type": "FloatDistribution",
+ "low": getattr(distribution, "low"),
+ "high": getattr(distribution, "high"),
+ "step": getattr(distribution, "step"),
+ "log": getattr(distribution, "log"),
+ }
+ return float_distribution
+ if distribution.__class__.__name__ == "UniformDistribution":
+ # Deprecated from Optuna v3.0
+ uniform: FloatDistributionJSON = {
+ "type": "FloatDistribution",
+ "low": getattr(distribution, "low"),
+ "high": getattr(distribution, "high"),
+ "step": 0,
+ "log": False,
+ }
+ return uniform
+ if distribution.__class__.__name__ == "LogUniformDistribution":
+ # Deprecated from Optuna v3.0
+ log_uniform: FloatDistributionJSON = {
+ "type": "FloatDistribution",
+ "low": getattr(distribution, "low"),
+ "high": getattr(distribution, "high"),
+ "step": 0,
+ "log": True,
+ }
+ return log_uniform
+ if distribution.__class__.__name__ == "DiscreteUniformDistribution":
+ # Deprecated from Optuna v3.0
+ discrete_uniform: FloatDistributionJSON = {
+ "type": "FloatDistribution",
+ "low": getattr(distribution, "low"),
+ "high": getattr(distribution, "high"),
+ "step": getattr(distribution, "q"),
+ "log": False,
+ }
+ return discrete_uniform
+ if distribution.__class__.__name__ == "IntDistribution":
+ # Added from Optuna v3.0
+ int_distribution: IntDistributionJSON = {
+ "type": "IntDistribution",
+ "low": getattr(distribution, "low"),
+ "high": getattr(distribution, "high"),
+ "step": getattr(distribution, "step"),
+ "log": getattr(distribution, "log"),
+ }
+ return int_distribution
+ if distribution.__class__.__name__ == "IntUniformDistribution":
+ # Deprecated from Optuna v3.0
+ int_uniform: IntDistributionJSON = {
+ "type": "IntDistribution",
+ "low": getattr(distribution, "low"),
+ "high": getattr(distribution, "high"),
+ "step": getattr(distribution, "step"),
+ "log": False,
+ }
+ return int_uniform
+ if distribution.__class__.__name__ == "IntLogUniformDistribution":
+ # Deprecated from Optuna v3.0
+ int_log_uniform: IntDistributionJSON = {
+ "type": "IntDistribution",
+ "low": getattr(distribution, "low"),
+ "high": getattr(distribution, "high"),
+ "step": getattr(distribution, "step"),
+ "log": True,
+ }
+ return int_log_uniform
+ if isinstance(distribution, CategoricalDistribution):
+ categorical: CategoricalDistributionJSON = {
+ "type": "CategoricalDistribution",
+ "choices": [
+ {"pytype": str(type(choice)), "value": str(choice)}
+ for choice in distribution.choices
+ ],
+ }
+ return categorical
+ raise ValueError(f"Unexpected distribution {str(distribution)}")
+
+
def serialize_search_space(
search_space: list[tuple[str, BaseDistribution]]
) -> list[dict[str, Any]]:
@@ -166,8 +302,7 @@ def serialize_search_space(
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/AppDrawer.tsx b/optuna_dashboard/ts/components/AppDrawer.tsx
index bc889c7e..99306435 100644
--- a/optuna_dashboard/ts/components/AppDrawer.tsx
+++ b/optuna_dashboard/ts/components/AppDrawer.tsx
@@ -33,6 +33,13 @@ import { Switch } from "@mui/material"
const drawerWidth = 240
+export type PageId =
+ | "history"
+ | "analytics"
+ | "trialTable"
+ | "trialList"
+ | "note"
+
const openedMixin = (theme: Theme): CSSObject => ({
width: drawerWidth,
transition: theme.transitions.create("width", {
diff --git a/optuna_dashboard/ts/components/BestTrialsCard.tsx b/optuna_dashboard/ts/components/BestTrialsCard.tsx
index c1a5d7b4..d34007e5 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 = [
@@ -88,26 +91,23 @@ export const BestTrialsCard: FC<{
URL_PREFIX +
`/studies/${trial.study_id}/trials?numbers=${trial.number}`
}
+ sx={{ flexDirection: "column", alignItems: "flex-start" }}
>
Trial {trial.number}
}
- secondary={
- <>
-
- Objective Values = [{trial.values?.join(", ")}]
-
-
- Params = [
- {trial.params
- .map((p) => `${p.name}: ${p.value}`)
- .join(", ")}
- ]
-
- >
- }
/>
+
+ Objective Values = [{trial.values?.join(", ")}]
+
+
+ Params = [
+ {trial.params
+ .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..11ccf617 100644
--- a/optuna_dashboard/ts/components/GraphContour.tsx
+++ b/optuna_dashboard/ts/components/GraphContour.tsx
@@ -12,6 +12,7 @@ import {
Box,
} from "@mui/material"
import { plotlyDarkTemplate } from "./PlotlyDarkMode"
+import { useMergedUnionSearchSpace } from "../searchSpace"
// eslint-disable-next-line @typescript-eslint/no-explicit-any
const unique = (array: any[]) => {
@@ -38,26 +39,28 @@ export const Contour: FC<{
}> = ({ 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 searchSpace = useMergedUnionSearchSpace(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 && searchSpace.length > 0) {
+ setXParam(searchSpace[0])
}
- if (!yParam && paramNames && paramNames.length > 1) {
- setYParam(paramNames[1])
+ if (yParam === null && searchSpace.length > 1) {
+ setYParam(searchSpace[1])
}
const handleObjectiveChange = (event: SelectChangeEvent) => {
setObjectiveId(event.target.value as number)
}
const handleXParamChange = (event: SelectChangeEvent) => {
- setXParam(event.target.value as string)
+ const param = searchSpace.find((s) => s.name === event.target.value)
+ setXParam(param || null)
}
const handleYParamChange = (event: SelectChangeEvent) => {
- setYParam(event.target.value as string)
+ const param = searchSpace.find((s) => s.name === event.target.value)
+ setYParam(param || null)
}
useEffect(() => {
@@ -66,7 +69,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 +101,7 @@ export const Contour: FC<{
x:
-
y:
-
+
{space.map((d, i) => (