This commit is contained in:
moririn2528 committed 2023-09-07 19:00:38 +09:00
commit b10fd8be16
24 files changed
+357 -86

No files matched your search

+1 -1
View File
@@ -45,4 +45,4 @@ jobs:
with:
token: ${{ secrets.CODECOV_TOKEN }}
file: ./coverage.xml
fail_ci_if_error: true
fail_ci_if_error: false
+1
View File
@@ -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()
+2 -1
View File
@@ -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"
+4
View File
@@ -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")
+134
View File
@@ -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
+7 -2
View File
@@ -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
+26 -36
View File
@@ -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))
+15 -9
View File
@@ -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]
+1 -1
View File
@@ -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 {}
+2
View File
@@ -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,
}
})
}
+5 -2
View File
@@ -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" }} />
}
+6
View File
@@ -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 = {
+1
View File
@@ -43,6 +43,7 @@ docs = [
test = [
"coverage",
"plotly",
"pytest",
"moto[s3]",
]
+1 -5
View File
@@ -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
View File
@@ -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 [
+86
View File
@@ -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
+2 -2
View File
@@ -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"]