mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-10-03 12:41:13 +08:00
merge
This commit is contained in:
commit
b10fd8be16
24 files changed
+357
-86
No files matched your search
@@ -45,4 +45,4 @@ jobs:
|
||||
with:
|
||||
token: ${{ secrets.CODECOV_TOKEN }}
|
||||
file: ./coverage.xml
|
||||
fail_ci_if_error: true
|
||||
fail_ci_if_error: false
|
||||
@@ -14,6 +14,7 @@ General APIs
|
||||
optuna_dashboard.wsgi
|
||||
optuna_dashboard.set_objective_names
|
||||
optuna_dashboard.save_note
|
||||
optuna_dashboard.save_plotly_graph_object
|
||||
|
||||
Human-in-the-loop
|
||||
-----------------
|
||||
|
||||
@@ -64,9 +64,6 @@ def main() -> NoReturn:
|
||||
)
|
||||
save_note(trial, note)
|
||||
|
||||
# 5. Mark comparison ready
|
||||
study.mark_comparison_ready(trial)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,5 +1,6 @@
|
||||
from ._app import run_server # noqa
|
||||
from ._app import wsgi # noqa
|
||||
from ._custom_plot_data import save_plotly_graph_object # noqa
|
||||
from ._form_widget import ChoiceWidget # noqa
|
||||
from ._form_widget import dict_to_form_widget # noqa
|
||||
from ._form_widget import ObjectiveChoiceWidget # noqa
|
||||
@@ -15,4 +16,4 @@ from ._note import get_note # noqa
|
||||
from ._note import save_note # noqa
|
||||
|
||||
|
||||
__version__ = "0.12.0"
|
||||
__version__ = "0.13.0b1"
|
||||
@@ -25,6 +25,7 @@ from . import _note as note
|
||||
from ._bottle_util import BottleViewReturn
|
||||
from ._bottle_util import json_api_view
|
||||
from ._cached_extra_study_property import get_cached_extra_study_property
|
||||
from ._custom_plot_data import get_plotly_graph_objects
|
||||
from ._importance import get_param_importance_from_trials_cache
|
||||
from ._pareto_front import get_pareto_front_trials
|
||||
from ._preferential_history import NewHistory
|
||||
@@ -214,6 +215,8 @@ def create_app(
|
||||
union_user_attrs,
|
||||
has_intermediate_values,
|
||||
) = get_cached_extra_study_property(study_id, trials)
|
||||
|
||||
plotly_graph_objects = get_plotly_graph_objects(system_attrs)
|
||||
return serialize_study_detail(
|
||||
summary,
|
||||
best_trials,
|
||||
@@ -222,6 +225,7 @@ def create_app(
|
||||
union,
|
||||
union_user_attrs,
|
||||
has_intermediate_values,
|
||||
plotly_graph_objects,
|
||||
)
|
||||
|
||||
@app.get("/api/studies/<study_id:int>/param_importances")
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from typing import TYPE_CHECKING
|
||||
import uuid
|
||||
|
||||
from optuna import Study
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from typing import Any
|
||||
|
||||
from optuna.storages import BaseStorage
|
||||
import plotly.graph_objs as go
|
||||
|
||||
|
||||
SYSTEM_ATTR_PLOT_DATA = "dashboard:plot_data:"
|
||||
SYSTEM_ATTR_MAX_LENGTH = 2045
|
||||
|
||||
|
||||
def save_plotly_graph_object(
|
||||
study: Study, figure: go.Figure, *, graph_object_id: str | None = None
|
||||
) -> str:
|
||||
"""Save the user-defined plotly's graph object to the study.
|
||||
|
||||
Example:
|
||||
|
||||
.. code-block:: python
|
||||
|
||||
import optuna
|
||||
from optuna_dashboard import save_plotly_graph_object
|
||||
|
||||
def objective(trial):
|
||||
x = trial.suggest_float("x", -100, 100)
|
||||
y = trial.suggest_categorical("y", [-1, 0, 1])
|
||||
return x**2 + y
|
||||
|
||||
study = optuna.create_study()
|
||||
study.optimize(objective, n_trials=100)
|
||||
|
||||
figure = optuna.visualization.plot_optimization_history(study)
|
||||
save_plotly_graph_object(study, figure)
|
||||
|
||||
Args:
|
||||
study:
|
||||
Target study object.
|
||||
plot_data:
|
||||
The plotly's graph object to save.
|
||||
graph_object_id:
|
||||
Unique identifier of the graph object. If specified, the graph object is overwritten.
|
||||
This must be a valid HTML id attribute value.
|
||||
|
||||
Returns:
|
||||
The graph object ID.
|
||||
"""
|
||||
if graph_object_id is not None and not is_valid_graph_object_id(graph_object_id):
|
||||
raise ValueError("graph_object_id must be a valid HTML id attribute value.")
|
||||
|
||||
storage = study._storage
|
||||
study_id = study._study_id
|
||||
|
||||
graph_object_id = graph_object_id or str(uuid.uuid4())
|
||||
key = SYSTEM_ATTR_PLOT_DATA + graph_object_id + ":"
|
||||
plot_data_json_str = figure.to_json()
|
||||
save_graph_object_json(storage, study_id, key, plot_data_json_str)
|
||||
return graph_object_id
|
||||
|
||||
|
||||
def save_graph_object_json(
|
||||
storage: BaseStorage, study_id: int, key_prefix: str, plot_data_json_str: str
|
||||
) -> None:
|
||||
plot_data_system_attrs = split_plot_data(plot_data_json_str, key_prefix)
|
||||
for k, v in plot_data_system_attrs.items():
|
||||
storage.set_study_system_attr(study_id, k, v)
|
||||
|
||||
# Clear previous graph object attributes
|
||||
study_system_attrs = storage.get_study_system_attrs(study_id)
|
||||
all_plot_data_system_attrs = [k for k in study_system_attrs if k.startswith(key_prefix)]
|
||||
if len(all_plot_data_system_attrs) > len(plot_data_system_attrs):
|
||||
for i in range(len(plot_data_system_attrs), len(all_plot_data_system_attrs)):
|
||||
storage.set_study_system_attr(study_id, f"{key_prefix}{i}", "")
|
||||
|
||||
|
||||
def list_graph_object_ids(system_attrs: dict[str, Any]) -> list[str]:
|
||||
titles = set()
|
||||
for key in system_attrs:
|
||||
if not key.startswith(SYSTEM_ATTR_PLOT_DATA):
|
||||
continue
|
||||
|
||||
s = key.split(":", maxsplit=2) # e.g. ["dashboard", "plot_data", "Optimization History:1"]
|
||||
if len(s) != 3:
|
||||
continue
|
||||
# Please note that title may contain ":".
|
||||
title = s[2].rsplit(":", maxsplit=1)[0]
|
||||
titles.add(title)
|
||||
return list(titles)
|
||||
|
||||
|
||||
def get_plotly_graph_objects(system_attrs: dict[str, Any]) -> dict[str, str]:
|
||||
graph_objects = {}
|
||||
for title in list_graph_object_ids(system_attrs):
|
||||
key_prefix = SYSTEM_ATTR_PLOT_DATA + title + ":"
|
||||
plot_data_attrs = {k: v for k, v in system_attrs.items() if k.startswith(key_prefix)}
|
||||
graph_objects[title] = concat_plot_data(plot_data_attrs, key_prefix)
|
||||
return graph_objects
|
||||
|
||||
|
||||
def split_plot_data(plot_data_str: str, key_prefix: str) -> dict[str, str]:
|
||||
plot_data_len = len(plot_data_str)
|
||||
attrs = {}
|
||||
for i in range(math.ceil(plot_data_len / SYSTEM_ATTR_MAX_LENGTH)):
|
||||
start = i * SYSTEM_ATTR_MAX_LENGTH
|
||||
end = min((i + 1) * SYSTEM_ATTR_MAX_LENGTH, plot_data_len)
|
||||
attrs[f"{key_prefix}{i}"] = plot_data_str[start:end]
|
||||
return attrs
|
||||
|
||||
|
||||
def concat_plot_data(plot_data_attrs: dict[str, str], key_prefix: str) -> str:
|
||||
return "".join(plot_data_attrs[f"{key_prefix}{i}"] for i in range(len(plot_data_attrs)))
|
||||
|
||||
|
||||
def is_valid_graph_object_id(graph_object_id: str) -> bool:
|
||||
if len(graph_object_id) == 0:
|
||||
return False
|
||||
|
||||
# Can only contain letters [A-Za-z], numbers [0-9], hyphens ("-"), underscores ("_"),
|
||||
# colons, and periods.
|
||||
if not all(
|
||||
"a" <= c <= "z" or "A" <= c <= "Z" or "0" <= c <= "9" or c in ("-", "_", ":", ".")
|
||||
for c in graph_object_id[1:]
|
||||
):
|
||||
return False
|
||||
# Unlike HTML id attribute, graph object id can begin with a letter [A-Za-z]
|
||||
return True
|
||||
@@ -18,7 +18,7 @@ from ._named_objectives import get_objective_names
|
||||
from ._preferential_history import _SYSTEM_ATTR_PREFIX_HISTORY
|
||||
from .artifact._backend import list_trial_artifacts
|
||||
from .preferential._study import _SYSTEM_ATTR_PREFERENTIAL_STUDY
|
||||
from .preferential._system_attrs import _get_preferences
|
||||
from .preferential._system_attrs import get_preferences
|
||||
|
||||
|
||||
if TYPE_CHECKING:
|
||||
@@ -133,6 +133,7 @@ def serialize_study_detail(
|
||||
union: list[tuple[str, BaseDistribution]],
|
||||
union_user_attrs: list[tuple[str, bool]],
|
||||
has_intermediate_values: bool,
|
||||
plotly_graph_objects: dict[str, str],
|
||||
) -> dict[str, Any]:
|
||||
serialized: dict[str, Any] = {
|
||||
"name": summary.study_name,
|
||||
@@ -163,7 +164,11 @@ def serialize_study_detail(
|
||||
serialized["form_widgets"] = form_widgets
|
||||
if serialized["is_preferential"]:
|
||||
serialized["preference_history"] = serialize_preference_history(system_attrs)
|
||||
serialized["preferences"] = _get_preferences(system_attrs)
|
||||
serialized["preferences"] = get_preferences(system_attrs)
|
||||
serialized["plotly_graph_objects"] = [
|
||||
{"id": id_, "graph_object": graph_object}
|
||||
for id_, graph_object in plotly_graph_objects.items()
|
||||
]
|
||||
return serialized
|
||||
|
||||
|
||||
|
||||
@@ -14,6 +14,7 @@ from optuna.trial import FrozenTrial
|
||||
from optuna.trial import TrialState
|
||||
from optuna_dashboard.preferential._system_attrs import get_n_generate
|
||||
from optuna_dashboard.preferential._system_attrs import get_preferences
|
||||
from optuna_dashboard.preferential._system_attrs import get_skipped_trial_ids
|
||||
from optuna_dashboard.preferential._system_attrs import is_skipped_trial
|
||||
from optuna_dashboard.preferential._system_attrs import report_preferences
|
||||
from optuna_dashboard.preferential._system_attrs import set_n_generate
|
||||
@@ -21,7 +22,6 @@ from optuna_dashboard.preferential._system_attrs import set_n_generate
|
||||
|
||||
_logger = logging.get_logger(__name__)
|
||||
_SYSTEM_ATTR_PREFERENTIAL_STUDY = "preference:is_preferential"
|
||||
_SYSTEM_ATTR_COMPARISON_READY = "preference:comparison_ready"
|
||||
|
||||
|
||||
class PreferentialStudy:
|
||||
@@ -62,13 +62,6 @@ class PreferentialStudy:
|
||||
def best_trials(self) -> list[FrozenTrial]:
|
||||
"""Return the trials that is not dominated by other trials.
|
||||
|
||||
.. seealso::
|
||||
|
||||
See `Study.best_trials`_ for details.
|
||||
|
||||
.. _Study.best_trials: https://optuna.readthedocs.io/en/stable/reference/\
|
||||
generated/optuna.study.Study.html#optuna.study.Study.best_trials
|
||||
|
||||
Returns:
|
||||
A list of FrozenTrial object
|
||||
"""
|
||||
@@ -251,8 +244,11 @@ class PreferentialStudy:
|
||||
Returns:
|
||||
A list of the pair of FrozenTrial objects. The left trial is better than the right one.
|
||||
"""
|
||||
|
||||
preferences = get_preferences(
|
||||
self._study._storage.get_study_system_attrs(self._study._study_id)
|
||||
) # Must come before study.get_trials()
|
||||
trials = self._study.get_trials(deepcopy=deepcopy)
|
||||
preferences = get_preferences(self._study._study_id, self._study._storage)
|
||||
return [(trials[better], trials[worse]) for (better, worse) in preferences]
|
||||
|
||||
def set_user_attr(self, key: str, value: Any) -> None:
|
||||
@@ -269,24 +265,6 @@ class PreferentialStudy:
|
||||
"""
|
||||
self._study.set_user_attr(key, value)
|
||||
|
||||
def mark_comparison_ready(self, trial_or_number: optuna.Trial | int) -> None:
|
||||
"""Mark trials ready to compare.
|
||||
|
||||
Args:
|
||||
trial_or_number:
|
||||
A Trial object or trial_number.
|
||||
"""
|
||||
storage = self._study._storage
|
||||
if isinstance(trial_or_number, optuna.Trial):
|
||||
trial_id = trial_or_number._trial_id
|
||||
elif isinstance(trial_or_number, int):
|
||||
trial_id = storage.get_trial_id_from_study_id_trial_number(
|
||||
self._study._study_id, trial_or_number
|
||||
)
|
||||
else:
|
||||
raise RuntimeError("Unexpected trial type")
|
||||
storage.set_trial_system_attr(trial_id, _SYSTEM_ATTR_COMPARISON_READY, True)
|
||||
|
||||
def should_generate(self) -> bool:
|
||||
"""Return whether the generator should generate a new trial now.
|
||||
|
||||
@@ -295,21 +273,33 @@ class PreferentialStudy:
|
||||
to generate a new trial if this method returns :obj:`True`, and to wait for human
|
||||
evaluation if this method returns :obj:`False`.
|
||||
"""
|
||||
return len(self.best_trials) < get_n_generate(self._study.system_attrs)
|
||||
study_system_attrs = self._study._storage.get_study_system_attrs(
|
||||
self._study._study_id
|
||||
) # Must come before _study.get_trials()
|
||||
trials = self._study.get_trials(
|
||||
deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING)
|
||||
)
|
||||
worse_trial_numbers = {worse for _, worse in get_preferences(study_system_attrs)}
|
||||
skipped_trial_ids = set(get_skipped_trial_ids(study_system_attrs))
|
||||
active_trials = [
|
||||
t
|
||||
for t in trials
|
||||
if t.number not in worse_trial_numbers and t._trial_id not in skipped_trial_ids
|
||||
]
|
||||
return len(active_trials) < get_n_generate(self._study.system_attrs)
|
||||
|
||||
|
||||
def get_best_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]:
|
||||
preferences = get_preferences(study_id, storage)
|
||||
preferences = get_preferences(storage.get_study_system_attrs(study_id))
|
||||
worse_numbers = {worse for _, worse in preferences}
|
||||
nondominated_numbers = {better for better, _ in preferences if better not in worse_numbers}
|
||||
trials = storage.get_all_trials(study_id, deepcopy=False)
|
||||
|
||||
study_system_attrs = storage.get_study_system_attrs(study_id)
|
||||
|
||||
best_trials = []
|
||||
for t in storage.get_all_trials(
|
||||
study_id, deepcopy=False, states=(TrialState.COMPLETE, TrialState.RUNNING)
|
||||
):
|
||||
if not t.system_attrs.get(_SYSTEM_ATTR_COMPARISON_READY, False):
|
||||
continue
|
||||
if t.number in worse_numbers:
|
||||
continue
|
||||
for n in nondominated_numbers:
|
||||
t = trials[n]
|
||||
if is_skipped_trial(t._trial_id, study_system_attrs):
|
||||
continue
|
||||
best_trials.append(copy.deepcopy(t))
|
||||
|
||||
@@ -35,22 +35,15 @@ def report_preferences(
|
||||
return preference_id
|
||||
|
||||
|
||||
def _get_preferences(system_attrs: dict[str, Any]) -> list[tuple[int, int]]:
|
||||
def get_preferences(study_system_attrs: dict[str, Any]) -> list[tuple[int, int]]:
|
||||
preferences: list[tuple[int, int]] = []
|
||||
for k, v in system_attrs.items():
|
||||
for k, v in study_system_attrs.items():
|
||||
if not k.startswith(_SYSTEM_ATTR_PREFIX_PREFERENCE):
|
||||
continue
|
||||
preferences.extend(v) # type: ignore
|
||||
return preferences
|
||||
|
||||
|
||||
def get_preferences(
|
||||
study_id: int,
|
||||
storage: BaseStorage,
|
||||
) -> list[tuple[int, int]]:
|
||||
return _get_preferences(storage.get_study_system_attrs(study_id))
|
||||
|
||||
|
||||
def report_skip(
|
||||
study_id: int,
|
||||
trial_id: int,
|
||||
@@ -68,6 +61,19 @@ def is_skipped_trial(trial_id: int, study_system_attrs: dict[str, Any]) -> bool:
|
||||
return key in study_system_attrs
|
||||
|
||||
|
||||
def get_skipped_trial_ids(study_system_attrs: dict[str, Any]) -> list[int]:
|
||||
skipped_trial_ids: list[int] = []
|
||||
for k in study_system_attrs:
|
||||
if not k.startswith(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL):
|
||||
continue
|
||||
try:
|
||||
trial_id = int(k[len(_SYSTEM_ATTR_PREFIX_SKIP_TRIAL) :]) # noqa: E203
|
||||
skipped_trial_ids.append(trial_id)
|
||||
except ValueError:
|
||||
continue
|
||||
return skipped_trial_ids
|
||||
|
||||
|
||||
def get_n_generate(study_system_attrs: dict[str, Any]) -> int:
|
||||
return study_system_attrs[_SYSTEM_ATTR_N_GENERATE]
|
||||
|
||||
|
||||
@@ -342,7 +342,7 @@ class PreferentialGPSampler(optuna.samplers.BaseSampler):
|
||||
if len(search_space) == 0:
|
||||
return {}
|
||||
|
||||
preferences = get_preferences(study._study_id, study._storage)
|
||||
preferences = get_preferences(study.system_attrs)
|
||||
trials = study.get_trials(deepcopy=False)
|
||||
if len(preferences) == 0:
|
||||
return {}
|
||||
|
||||
@@ -94,6 +94,7 @@ interface StudyDetailResponse {
|
||||
form_widgets?: FormWidgets
|
||||
preferences?: [number, number][]
|
||||
preference_history?: PreferenceHistoryResponce[]
|
||||
plotly_graph_objects: PlotlyGraphObject[]
|
||||
}
|
||||
|
||||
export const getStudyDetailAPI = (
|
||||
@@ -133,6 +134,7 @@ export const getStudyDetailAPI = (
|
||||
preference_history: res.data.preference_history?.map(
|
||||
convertPreferenceHistory
|
||||
),
|
||||
plotly_graph_objects: res.data.plotly_graph_objects,
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
@@ -249,7 +249,7 @@ export const AppDrawer: FC<{
|
||||
</ListItemButton>
|
||||
</ListItem>
|
||||
{studyDetail?.is_preferential && (
|
||||
<ListItem key="Graph" disablePadding sx={styleListItem}>
|
||||
<ListItem key="PreferenceGraph" disablePadding sx={styleListItem}>
|
||||
<ListItemButton
|
||||
component={Link}
|
||||
to={`${URL_PREFIX}/studies/${studyId}/graph`}
|
||||
@@ -259,7 +259,10 @@ export const AppDrawer: FC<{
|
||||
<ListItemIcon sx={styleListItemIcon}>
|
||||
<LanIcon />
|
||||
</ListItemIcon>
|
||||
<ListItemText primary="Graph" sx={styleListItemText} />
|
||||
<ListItemText
|
||||
primary="PreferenceGraph"
|
||||
sx={styleListItemText}
|
||||
/>
|
||||
</ListItemButton>
|
||||
</ListItem>
|
||||
)}
|
||||
|
||||
@@ -101,7 +101,10 @@ function reductionPreference(
|
||||
let n = 0
|
||||
for (const [source, target] of input_preferences) {
|
||||
if (
|
||||
preferences.find((p) => p[0] === source && p[1] === target) !== undefined
|
||||
preferences.find((p) => p[0] === source && p[1] === target) !==
|
||||
undefined ||
|
||||
input_preferences.find((p) => p[0] === target && p[1] === source) !==
|
||||
undefined
|
||||
) {
|
||||
continue
|
||||
}
|
||||
@@ -139,7 +142,7 @@ function reductionPreference(
|
||||
}
|
||||
if (topologicalOrder.length !== n) {
|
||||
console.error("cycle detected")
|
||||
return []
|
||||
return preferences
|
||||
}
|
||||
|
||||
const response: [number, number][] = []
|
||||
|
||||
@@ -42,6 +42,8 @@ const PreferentialTrial: FC<{
|
||||
)
|
||||
}
|
||||
|
||||
const isBestTrial = trial.state === "Complete"
|
||||
|
||||
return (
|
||||
<Card
|
||||
sx={{
|
||||
@@ -159,7 +161,7 @@ const PreferentialTrial: FC<{
|
||||
>
|
||||
<TrialListDetail
|
||||
trial={trial}
|
||||
isBestTrial={() => true}
|
||||
isBestTrial={() => isBestTrial}
|
||||
directions={[]}
|
||||
objectiveNames={[]}
|
||||
/>
|
||||
@@ -182,11 +184,15 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({
|
||||
return null
|
||||
}
|
||||
const theme = useTheme()
|
||||
|
||||
const runningTrials = studyDetail.trials.filter((t) => t.state === "Running")
|
||||
const activeTrials = runningTrials.concat(studyDetail.best_trials)
|
||||
|
||||
const [displayTrials, setDisplayTrials] = useState<DisplayTrials>({
|
||||
numbers: studyDetail.best_trials.map((t) => t.number),
|
||||
last_number: Math.max(...studyDetail.best_trials.map((t) => t.number), -1),
|
||||
numbers: activeTrials.map((t) => t.number),
|
||||
last_number: Math.max(...activeTrials.map((t) => t.number), -1),
|
||||
})
|
||||
const new_trails = studyDetail.best_trials.filter(
|
||||
const new_trails = activeTrials.filter(
|
||||
(t) =>
|
||||
displayTrials.last_number < t.number &&
|
||||
displayTrials.numbers.find((n) => n === t.number) === undefined
|
||||
@@ -239,7 +245,7 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({
|
||||
{displayTrials.numbers.map((t, index) => (
|
||||
<PreferentialTrial
|
||||
key={index}
|
||||
trial={studyDetail.best_trials.find((trial) => trial.number === t)}
|
||||
trial={activeTrials.find((trial) => trial.number === t)}
|
||||
candidates={displayTrials.numbers.filter((n) => n !== -1)}
|
||||
hideTrial={() => {
|
||||
hideTrial(t)
|
||||
|
||||
@@ -15,6 +15,7 @@ import { GraphIntermediateValues } from "./GraphIntermediateValues"
|
||||
import Grid2 from "@mui/material/Unstable_Grid2"
|
||||
import { DataGrid, DataGridColumn } from "./DataGrid"
|
||||
import { GraphHyperparameterImportance } from "./GraphHyperparameterImportances"
|
||||
import { UserDefinedPlot } from "./UserDefinedPlot"
|
||||
import { BestTrialsCard } from "./BestTrialsCard"
|
||||
import {
|
||||
useStudyDetailValue,
|
||||
@@ -124,6 +125,16 @@ export const StudyHistory: FC<{ studyId: number }> = ({ studyId }) => {
|
||||
<Grid2 xs={6}>
|
||||
<GraphTimeline study={studyDetail} />
|
||||
</Grid2>
|
||||
{studyDetail !== null &&
|
||||
studyDetail.plotly_graph_objects.map((go) => (
|
||||
<Grid2 xs={6} key={go.id}>
|
||||
<Card>
|
||||
<CardContent>
|
||||
<UserDefinedPlot graphObject={go} />
|
||||
</CardContent>
|
||||
</Card>
|
||||
</Grid2>
|
||||
))}
|
||||
<Grid2 xs={6} spacing={2}>
|
||||
<BestTrialsCard studyDetail={studyDetail} />
|
||||
</Grid2>
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
import * as plotly from "plotly.js-dist-min"
|
||||
import React, { FC, useEffect } from "react"
|
||||
import { Box } from "@mui/material"
|
||||
|
||||
export const UserDefinedPlot: FC<{
|
||||
graphObject: PlotlyGraphObject
|
||||
}> = ({ graphObject }) => {
|
||||
const plotDomId = `user-defined-plot:${graphObject.id}`
|
||||
|
||||
useEffect(() => {
|
||||
try {
|
||||
const parsed = JSON.parse(graphObject.graph_object)
|
||||
plotly.react(plotDomId, parsed.data, parsed.layout)
|
||||
} catch (e) {
|
||||
// Avoid to crash the whole page when given invalid grpah objects.
|
||||
console.error(e)
|
||||
}
|
||||
}, [graphObject])
|
||||
|
||||
return <Box id={plotDomId} sx={{ height: "450px" }} />
|
||||
}
|
||||
Vendored
+6
@@ -182,6 +182,11 @@ type FormWidgets =
|
||||
widgets: UserAttrFormWidget[]
|
||||
}
|
||||
|
||||
type PlotlyGraphObject = {
|
||||
id: string
|
||||
graph_object: string
|
||||
}
|
||||
|
||||
type StudyDetail = {
|
||||
id: number
|
||||
name: string
|
||||
@@ -200,6 +205,7 @@ type StudyDetail = {
|
||||
form_widgets?: FormWidgets
|
||||
preferences?: [number, number][]
|
||||
preference_history?: PreferenceHistory[]
|
||||
plotly_graph_objects: PlotlyGraphObject[]
|
||||
}
|
||||
|
||||
type StudyDetails = {
|
||||
|
||||
@@ -43,6 +43,7 @@ docs = [
|
||||
|
||||
test = [
|
||||
"coverage",
|
||||
"plotly",
|
||||
"pytest",
|
||||
"moto[s3]",
|
||||
]
|
||||
|
||||
@@ -40,7 +40,6 @@ def test_report_and_get_preferences(storage_supplier: Callable[[], StorageSuppli
|
||||
for _ in range(2):
|
||||
trial = study.ask()
|
||||
trial.suggest_float("x", 0, 1)
|
||||
study.mark_comparison_ready(trial)
|
||||
better, worse = study.trials
|
||||
study.report_preference(better, worse)
|
||||
assert len(study.preferences) == 1
|
||||
@@ -152,7 +151,6 @@ def test_copy_study() -> None:
|
||||
for _ in range(3):
|
||||
trial = from_study.ask()
|
||||
trial.suggest_float("x", 0, 1)
|
||||
from_study.mark_comparison_ready(trial)
|
||||
from_study.report_preference(from_study.trials[0], from_study.trials[1])
|
||||
from_study.report_preference(from_study.trials[1], from_study.trials[2])
|
||||
|
||||
@@ -243,7 +241,6 @@ def test_get_trials(storage_supplier: Callable[[], StorageSupplier]) -> None:
|
||||
for _ in range(5):
|
||||
trial = study.ask()
|
||||
trial.suggest_int("x", 1, 5)
|
||||
study.mark_comparison_ready(trial)
|
||||
|
||||
with patch("copy.deepcopy", wraps=copy.deepcopy) as mock_object:
|
||||
trials0 = study.get_trials(deepcopy=False)
|
||||
@@ -266,8 +263,7 @@ def test_get_trials_state_option(storage_supplier: Callable[[], StorageSupplier]
|
||||
with storage_supplier() as storage:
|
||||
study = create_study(n_generate=4, storage=storage)
|
||||
for _ in range(3):
|
||||
trial = study.ask()
|
||||
study.mark_comparison_ready(trial)
|
||||
study.ask()
|
||||
better, worse = study.trials[:2]
|
||||
study.report_preference(better, worse)
|
||||
|
||||
|
||||
@@ -18,12 +18,13 @@ def test_report_and_get_preferences(storage_supplier: Callable[[], StorageSuppli
|
||||
study.ask()
|
||||
|
||||
study_id = study._study_id
|
||||
assert len(get_preferences(study_id, storage)) == 0
|
||||
|
||||
assert len(get_preferences(storage.get_study_system_attrs(study_id))) == 0
|
||||
|
||||
better, worse = study.trials[0], study.trials[1]
|
||||
report_preferences(study_id, storage, [(better.number, worse.number)])
|
||||
assert len(get_preferences(study_id, storage)) == 1
|
||||
assert len(get_preferences(storage.get_study_system_attrs(study_id))) == 1
|
||||
|
||||
actual_better, actual_worse = get_preferences(study_id, storage)[0]
|
||||
actual_better, actual_worse = get_preferences(storage.get_study_system_attrs(study_id))[0]
|
||||
assert actual_better == better.number
|
||||
assert actual_worse == worse.number
|
||||
+11
-13
@@ -104,10 +104,11 @@ class APITestCase(TestCase):
|
||||
storage = optuna.storages.InMemoryStorage()
|
||||
study = create_study(n_generate=4, storage=storage)
|
||||
for _ in range(3):
|
||||
trial = study.ask()
|
||||
study.mark_comparison_ready(trial)
|
||||
study.ask()
|
||||
study.report_preference(study.trials[0], study.trials[1])
|
||||
|
||||
assert len(study.best_trials) == 1
|
||||
|
||||
app = create_app(storage)
|
||||
study_id = study._study._study_id
|
||||
status, _, body = send_request(
|
||||
@@ -119,16 +120,14 @@ class APITestCase(TestCase):
|
||||
self.assertEqual(status, 200)
|
||||
|
||||
best_trials = json.loads(body)["best_trials"]
|
||||
assert len(best_trials) == 2
|
||||
assert len(best_trials) == 1
|
||||
assert best_trials[0]["number"] == 0
|
||||
assert best_trials[1]["number"] == 2
|
||||
|
||||
def test_report_preference(self) -> None:
|
||||
storage = optuna.storages.InMemoryStorage()
|
||||
study = create_study(n_generate=4, storage=storage)
|
||||
for _ in range(3):
|
||||
trial = study.ask()
|
||||
study.mark_comparison_ready(trial)
|
||||
study.ask()
|
||||
|
||||
app = create_app(storage)
|
||||
study_id = study._study._study_id
|
||||
@@ -161,8 +160,7 @@ class APITestCase(TestCase):
|
||||
storage = optuna.storages.InMemoryStorage()
|
||||
study = create_study(storage=storage, n_generate=3)
|
||||
for _ in range(3):
|
||||
trial = study.ask()
|
||||
study.mark_comparison_ready(trial)
|
||||
study.ask()
|
||||
|
||||
app = create_app(storage)
|
||||
study_id = study._study._study_id
|
||||
@@ -187,23 +185,23 @@ class APITestCase(TestCase):
|
||||
trials: list[optuna.Trial] = []
|
||||
for _ in range(3):
|
||||
trial = study.ask()
|
||||
study.mark_comparison_ready(trial)
|
||||
trials.append(trial)
|
||||
study.report_preference(study.trials[0], study.trials[1])
|
||||
study.report_preference(study.trials[2], study.trials[1])
|
||||
|
||||
app = create_app(storage)
|
||||
study_id = study._study._study_id
|
||||
status, _, _ = send_request(
|
||||
app,
|
||||
f"/api/studies/{study_id}/{trials[1]._trial_id}/skip",
|
||||
f"/api/studies/{study_id}/{trials[0]._trial_id}/skip",
|
||||
"POST",
|
||||
content_type="application/json",
|
||||
)
|
||||
self.assertEqual(status, 204)
|
||||
|
||||
best_trials = study.best_trials
|
||||
assert len(best_trials) == 2
|
||||
assert best_trials[0].number == 0
|
||||
assert best_trials[1].number == 2
|
||||
assert len(best_trials) == 1
|
||||
assert best_trials[0].number == 2
|
||||
|
||||
def test_create_study(self) -> None:
|
||||
for name, directions, expected_status in [
|
||||
|
||||
@@ -0,0 +1,86 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import optuna
|
||||
from optuna_dashboard import _custom_plot_data as custom_plot_data
|
||||
from optuna_dashboard import save_plotly_graph_object
|
||||
import pytest
|
||||
|
||||
|
||||
def get_dummy_study() -> optuna.Study:
|
||||
def objective(trial: optuna.Trial) -> float:
|
||||
x = trial.suggest_float("x", -100, 100)
|
||||
y = trial.suggest_categorical("y", [-1, 0, 1])
|
||||
return x**2 + y
|
||||
|
||||
study = optuna.create_study()
|
||||
optuna.logging.set_verbosity(optuna.logging.ERROR)
|
||||
study.optimize(objective, n_trials=100)
|
||||
return study
|
||||
|
||||
|
||||
def test_save_plotly_graph_object() -> None:
|
||||
# Save history plot
|
||||
dummy_study = get_dummy_study()
|
||||
plot_data = optuna.visualization.plot_optimization_history(dummy_study)
|
||||
graph_object_id = save_plotly_graph_object(dummy_study, plot_data)
|
||||
|
||||
study_system_attrs = dummy_study._storage.get_study_system_attrs(dummy_study._study_id)
|
||||
plot_data_dict = custom_plot_data.get_plotly_graph_objects(study_system_attrs)
|
||||
assert len(plot_data_dict) == 1
|
||||
assert plot_data_dict[graph_object_id] == plot_data.to_json()
|
||||
|
||||
# Save parallel coordinate plot
|
||||
plot_data = optuna.visualization.plot_parallel_coordinate(dummy_study)
|
||||
graph_object_id = save_plotly_graph_object(dummy_study, plot_data)
|
||||
|
||||
study_system_attrs = dummy_study._storage.get_study_system_attrs(dummy_study._study_id)
|
||||
plot_data_dict = custom_plot_data.get_plotly_graph_objects(study_system_attrs)
|
||||
assert len(plot_data_dict) == 2
|
||||
assert plot_data_dict[graph_object_id] == plot_data.to_json()
|
||||
|
||||
|
||||
def test_update_plotly_graph_object() -> None:
|
||||
# Save history plot
|
||||
dummy_study = get_dummy_study()
|
||||
plot_data = optuna.visualization.plot_optimization_history(dummy_study)
|
||||
graph_object_id = save_plotly_graph_object(dummy_study, plot_data)
|
||||
|
||||
study_system_attrs = dummy_study._storage.get_study_system_attrs(dummy_study._study_id)
|
||||
plot_data_dict = custom_plot_data.get_plotly_graph_objects(study_system_attrs)
|
||||
assert len(plot_data_dict) == 1
|
||||
assert plot_data_dict[graph_object_id] == plot_data.to_json()
|
||||
|
||||
# Save parallel coordinate plot
|
||||
plot_data = optuna.visualization.plot_parallel_coordinate(dummy_study)
|
||||
graph_object_id = save_plotly_graph_object(
|
||||
dummy_study, plot_data, graph_object_id=graph_object_id
|
||||
)
|
||||
|
||||
study_system_attrs = dummy_study._storage.get_study_system_attrs(dummy_study._study_id)
|
||||
plot_data_dict = custom_plot_data.get_plotly_graph_objects(study_system_attrs)
|
||||
assert len(plot_data_dict) == 1
|
||||
assert plot_data_dict[graph_object_id] == plot_data.to_json()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name",
|
||||
[
|
||||
"0",
|
||||
"a",
|
||||
"a1-:_.",
|
||||
],
|
||||
)
|
||||
def test_is_valid_graph_object_id(name: str) -> None:
|
||||
assert custom_plot_data.is_valid_graph_object_id(name)
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"name",
|
||||
[
|
||||
"a,",
|
||||
"a b",
|
||||
"aあいうえお",
|
||||
],
|
||||
)
|
||||
def test_is_invalid_graph_object_id(name: str) -> None:
|
||||
assert not custom_plot_data.is_valid_graph_object_id(name)
|
||||
@@ -19,7 +19,6 @@ def test_report_and_get_choices(storage_supplier: Callable[[], StorageSupplier])
|
||||
for _ in range(5):
|
||||
trial = study.ask()
|
||||
trial.suggest_float("x", 0, 1)
|
||||
study.mark_comparison_ready(trial)
|
||||
|
||||
study_id = study._study._study_id
|
||||
|
||||
|
||||
@@ -29,7 +29,7 @@ def test_get_study_detail_is_preferential() -> None:
|
||||
assert len(study_summaries) == 1
|
||||
|
||||
study_summary = study_summaries[0]
|
||||
study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False)
|
||||
study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False, {})
|
||||
assert study_detail["is_preferential"]
|
||||
|
||||
|
||||
@@ -40,7 +40,7 @@ def test_get_study_detail_is_not_preferential() -> None:
|
||||
assert len(study_summaries) == 1
|
||||
|
||||
study_summary = study_summaries[0]
|
||||
study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False)
|
||||
study_detail = serialize_study_detail(study_summary, [], study.trials, [], [], [], False, {})
|
||||
assert not study_detail["is_preferential"]
|
||||
|
||||
|
||||
|
||||
Reference in new issue
Block a user