Merge branch 'main' of github.com:optuna/optuna-dashboard into preferential-gp2

This commit is contained in:
Contramundum committed 2023-08-30 10:56:50 +09:00
commit d189e9dd45
35 files changed
+9278 -8525

No files matched your search

+4 -4
View File
@@ -11,10 +11,10 @@ The repository is organized as follows:
.
├── optuna_dashboard/ # The Python package.
│ └── ts/ # TypeScript code for the Python package.
├── standalone_app/ # Standalone application that can be run in browser or within the WebView of the VSCode extension.
├── standalone_app/ # Standalone application that can be run in browser or within the WebView of the VS Code extension.
│ ├── browser_app_entry.tsx # Entry point for browser app, hosted on GitHub pages.
│ └── vscode_entry.tsx # Entry point for VSCode app, output placed under `vscode/assets`.
├── vscode/ # The VSCode extension.
│ └── vscode_entry.tsx # Entry point for VS Code app, output placed under `vscode/assets`.
├── vscode/ # The VS Code extension.
└── rustlib/ # Rust library exporting Wasm functions.
└── pkg/ # Output directory for rustlib, installed from package.json via `"./rustlib/pkg"`.
```
@@ -133,7 +133,7 @@ $ make serve-browser-app
Open http://127.0.0.1:9000/
## VSCode Extension
## VS Code Extension
```
$ npm i -g vsce
+7
View File
@@ -75,6 +75,13 @@ $ docker run -it --rm -p 8080:8080 ghcr.io/optuna/optuna-dashboard postgresql+ps
</details>
## VS Code Extension (Experimental)
You can install the VS Code extension via [Visual Studio Marketplace](https://marketplace.visualstudio.com/items?itemName=Optuna.optuna-dashboard#overview).
![vscode-extension](./docs/_static/vscode-extension.png)
Please right-click the SQLite3 files (`*.db` or `*.sqlite3`) in the VS Code file explorer and select the "Open in Optuna Dashboard" command from the dropdown menu.
## Features
Binary file not shown.

After

Width:  |  Height:  |  Size: 542 KiB

+14
View File
@@ -18,6 +18,9 @@ General APIs
Human-in-the-loop
-----------------
Form Widgets
~~~~~~~~~~~~
.. autosummary::
:toctree: _generated/
:nosignatures:
@@ -30,6 +33,17 @@ Human-in-the-loop
optuna_dashboard.TextInputWidget
optuna_dashboard.ObjectiveUserAttrRef
Preferential Optimization
~~~~~~~~~~~~~~~~~~~~~~~~~
.. autosummary::
:toctree: _generated/
:nosignatures:
optuna_dashboard.preferential.create_study
optuna_dashboard.preferential.load_study
optuna_dashboard.preferential.PreferentialStudy
Streamlit
-----------------
+11
View File
@@ -201,3 +201,14 @@ you can use ``google.colab.output()`` function as follows:
output.serve_kernel_port_as_window(port, path='/dashboard/')
Then please open http://localhost:8081/dashboard to browse.
VS Code (Experimental)
----------------------
You can install the VS Code extension via `Visual Studio Marketplace <https://marketplace.visualstudio.com/items?itemName=Optuna.optuna-dashboard#overview>`_.
.. image:: _static/vscode-extension.png
:alt: Screenshot for the VS Code Extension
:align: center
Please right-click the SQLite3 files (`*.db` or `*.sqlite3`) in the file explorer and select the "Open in Optuna Dashboard" command from the dropdown menu.
@@ -11,7 +11,7 @@ from optuna_dashboard.artifact import get_artifact_path
from optuna_dashboard.artifact import upload_artifact
from optuna_dashboard.artifact.file_system import FileSystemBackend
from optuna_dashboard.preferential import create_study
from optuna_dashboard.preferential.samplers._gp import PreferentialGPSampler
from optuna_dashboard.preferential.samplers.gp import PreferentialGPSampler
from PIL import Image
+21 -1
View File
@@ -41,6 +41,7 @@ from .artifact._backend_to_store import to_artifact_store
from .preferential._study import _SYSTEM_ATTR_PREFERENTIAL_STUDY
from .preferential._study import get_best_trials as get_best_preferential_trials
from .preferential._system_attrs import report_preferences
from .preferential._system_attrs import report_skip
if typing.TYPE_CHECKING:
@@ -140,7 +141,9 @@ def create_app(
return {"reason": f"study_id={study_id} is not found"}
try:
dst_study = optuna.create_study(storage=storage, study_name=dst_study_name)
dst_study = optuna.create_study(
storage=storage, study_name=dst_study_name, directions=src_study.directions
)
dst_study.add_trials(src_study.get_trials(deepcopy=False))
except DuplicatedStudyError:
response.status = 400 # Bad request
@@ -329,6 +332,23 @@ def create_app(
response.status = 204
return {}
@app.post("/api/studies/<study_id:int>/<trial_id:int>/skip")
@json_api_view
def skip_trial(study_id: int, trial_id: int) -> dict[str, Any]:
try:
system_attrs = storage.get_study_system_attrs(study_id)
except KeyError:
response.status = 404 # Not found
return {"reason": f"study_id={study_id} is not found"}
is_preferential = system_attrs.get(_SYSTEM_ATTR_PREFERENTIAL_STUDY, False)
if not is_preferential:
response.status = 400 # Bad request
return {"reason": "The study is not preferential."}
report_skip(study_id, trial_id, storage)
response.status = 204 # No content
return {}
@app.put("/api/studies/<study_id:int>/<trial_id:int>/note")
@json_api_view
def save_trial_note(study_id: int, trial_id: int) -> dict[str, Any]:
+1
View File
@@ -131,6 +131,7 @@ def serialize_study_detail(
serialized: dict[str, Any] = {
"name": summary.study_name,
"directions": [d.name.lower() for d in summary.directions],
"user_attrs": serialize_attrs(summary.user_attrs),
}
system_attrs = getattr(summary, "system_attrs", {})
if summary.datetime_start is not None:
+34 -19
View File
@@ -40,8 +40,8 @@ if TYPE_CHECKING:
},
)
ARTIFACTS_ATTR_PREFIX = "dashboard:artifacts:"
ARTIFACTS_ATTR_PREFIX = "artifacts:"
DASHBOARD_ARTIFACTS_ATTR_PREFIX = "dashboard:artifacts:"
DEFAULT_MIME_TYPE = "application/octet-stream"
BaseRequest.MEMFILE_MAX = int(
os.environ.get("OPTUNA_DASHBOARD_MEMFILE_MAX", 1024 * 1024 * 128)
@@ -81,6 +81,14 @@ def register_artifact_route(
@app.post("/api/artifacts/<study_id:int>/<trial_id:int>")
@json_api_view
def upload_artifact_api(study_id: int, trial_id: int) -> dict[str, Any]:
trial = storage.get_trial(trial_id)
if trial is None:
response.status = 400
return {"reason": "Invalid study_id or trial_id"}
elif trial.state.is_finished():
response.status = 400
return {"reason": "The trial is already finished."}
# TODO(c-bata): Use optuna.artifacts.upload_artifact()
if artifact_store is None:
response.status = 400 # Bad Request
@@ -102,14 +110,10 @@ def register_artifact_route(
"mimetype": mimetype or DEFAULT_MIME_TYPE,
"encoding": encoding,
}
attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id
storage.set_study_system_attr(study_id, attr_key, json.dumps(artifact))
attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id
storage.set_trial_system_attr(trial_id, attr_key, json.dumps(artifact))
response.status = 201
trial = storage.get_trial(trial_id)
if trial is None:
response.status = 400
return {"reason": "Invalid study_id or trial_id"}
return {
"artifact_id": artifact_id,
"artifacts": list_trial_artifacts(storage.get_study_system_attrs(study_id), trial),
@@ -123,8 +127,14 @@ def register_artifact_route(
return {"reason": "Cannot access to the artifacts."}
artifact_store.remove(artifact_id)
attr_key = _artifact_prefix(trial_id) + artifact_id
storage.set_study_system_attr(study_id, attr_key, json.dumps(None))
# The artifact's metadata is stored in one of the following two locations:
storage.set_study_system_attr(
study_id, _artifact_prefix(trial_id) + artifact_id, json.dumps(None)
)
storage.set_trial_system_attr(
trial_id, ARTIFACTS_ATTR_PREFIX + artifact_id, json.dumps(None)
)
response.status = 204
return {}
@@ -169,7 +179,6 @@ def upload_artifact(
filename = os.path.basename(file_path)
storage = trial.storage
trial_id = trial._trial_id
study_id = trial.study._study_id
artifact_id = str(uuid.uuid4())
guess_mimetype, guess_encoding = mimetypes.guess_type(filename)
artifact: ArtifactMeta = {
@@ -178,8 +187,8 @@ def upload_artifact(
"encoding": encoding or guess_encoding,
"filename": filename,
}
attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id
storage.set_study_system_attr(study_id, attr_key, json.dumps(artifact))
attr_key = ARTIFACTS_ATTR_PREFIX + artifact_id
storage.set_trial_system_attr(trial_id, attr_key, json.dumps(artifact))
with open(file_path, "rb") as f:
backend.write(artifact_id, f)
@@ -187,23 +196,27 @@ def upload_artifact(
def _artifact_prefix(trial_id: int) -> str:
return ARTIFACTS_ATTR_PREFIX + f"{trial_id}:"
return DASHBOARD_ARTIFACTS_ATTR_PREFIX + f"{trial_id}:"
def get_artifact_meta(
storage: BaseStorage, study_id: int, trial_id: int, artifact_id: str
) -> Optional[ArtifactMeta]:
study_system_attr = storage.get_study_system_attrs(study_id)
# Search study_system_attrs due to backward compatibility.
study_system_attrs = storage.get_study_system_attrs(study_id)
attr_key = _artifact_prefix(trial_id=trial_id) + artifact_id
artifact_meta = study_system_attr.get(attr_key)
artifact_meta = study_system_attrs.get(attr_key)
if artifact_meta is not None:
return json.loads(artifact_meta)
# Search trial_system_attrs. Note that artifacts uploaded via optuna.artifacts.upload_artifact
# have a different trial_system_attrs key prefix.
# See https://github.com/optuna/optuna/blob/f827582a8/optuna/artifacts/_upload.py#L71
trial_system_attrs = storage.get_trial_system_attrs(trial_id)
value = trial_system_attrs.get("artifacts:" + artifact_id)
value = trial_system_attrs.get(ARTIFACTS_ATTR_PREFIX + artifact_id)
if value is not None:
return json.loads(value)
return None
@@ -221,18 +234,20 @@ def delete_all_artifacts(backend: ArtifactStore, storage: BaseStorage, study_id:
def list_trial_artifacts(
study_system_attrs: dict[str, Any], trial: FrozenTrial
) -> list[ArtifactMeta]:
# Collect ArtifactMeta from study_system_attrs due to backward compatibility.
dashboard_artifact_metas = [
json.loads(value)
for key, value in study_system_attrs.items()
if key.startswith(_artifact_prefix(trial._trial_id))
]
# Collect ArtifactMeta from trial_system_attrs. Note that artifacts uploaded via
# optuna.artifacts.upload_artifacts have a different trial_system_attrs key prefix.
# See https://github.com/optuna/optuna/blob/f827582a8/optuna/artifacts/_upload.py#L16
optuna_artifact_metas = [
json.loads(value)
for key, value in trial.system_attrs.items()
if key.startswith("artifacts:")
if key.startswith(ARTIFACTS_ATTR_PREFIX)
]
artifact_metas = dashboard_artifact_metas + optuna_artifact_metas
return [a for a in artifact_metas if a is not None]
+255 -10
View File
@@ -13,6 +13,7 @@ from optuna.samplers import RandomSampler
from optuna.trial import FrozenTrial
from optuna.trial import TrialState
from optuna_dashboard.preferential._system_attrs import get_preferences
from optuna_dashboard.preferential._system_attrs import is_skipped_trial
from optuna_dashboard.preferential._system_attrs import report_preferences
@@ -22,27 +23,87 @@ _SYSTEM_ATTR_COMPARISON_READY = "preference:comparison_ready"
class PreferentialStudy:
"""A Study-like class for preferential optimization.
This object provides interfaces to create a new `Trial`_, set/get results
of pairwise comparison called preferences.
.. _Trial: https://optuna.readthedocs.io/en/stable/reference/generated/\
optuna.trial.Trial.html#optuna.trial.Trial
Note that the direct use of this constructor is not recommended.
To create and load a study, please refer to the documentation of
:func:`~optuna_dashboard.preferential.create_study` and
:func:`~optuna_dashboard.preferential.load_study` respectively.
"""
def __init__(self, study: optuna.Study) -> None:
self._study = study
@property
def trials(self) -> list[FrozenTrial]:
"""Return the all trials.
.. seealso::
See `Study.trials`_ for details.
.. _Study.trials: https://optuna.readthedocs.io/en/stable/reference/generated/\
optuna.study.Study.html#optuna.study.Study.trials
Returns:
A list of FrozenTrial object
"""
return self._study.trials
@property
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
"""
return get_best_trials(self._study._study_id, self._study._storage)
@property
def study_name(self) -> str:
"""Return the name of the study.
Returns:
A string object
"""
return self._study.study_name
@property
def user_attrs(self) -> dict[str, Any]:
"""Return user attributes of the study.
.. seealso::
See `Study.user_attrs`_ for details.
.. _Study.user_attrs: https://optuna.readthedocs.io/en/stable/reference/\
generated/optuna.study.Study.html#optuna.study.Study.user_attrs
Returns:
A dictionary containing all user attributes
"""
return self._study.user_attrs
@property
def preferences(self) -> list[tuple[FrozenTrial, FrozenTrial]]:
"""Return results of pairwise comparison.
Returns:
A list of the pair of FrozenTrial objects. The left trial is better than the right one.
"""
return self.get_preferences(deepcopy=True)
def get_trials(
@@ -50,15 +111,73 @@ class PreferentialStudy:
deepcopy: bool = True,
states: Container[optuna.trial.TrialState] | None = None,
) -> list[FrozenTrial]:
"""Return the trials that is not dominated by other trials.
.. seealso::
See `Study.get_trials`_ for details.
.. _Study.get_trials: https://optuna.readthedocs.io/en/stable/reference/\
generated/optuna.study.Study.html#optuna.study.Study.get_trials
Args:
deepcopy:
Flag to control whether to apply ``copy.deepcopy()`` to the trials.
Note that if you set the flag to :obj:`False`, you shouldn't mutate
any fields of the returned trial. Otherwise the internal state of
the study may corrupt and unexpected behavior may happen.
states:
Trial states to filter on. If :obj:`None`, include all states.
Returns:
A list of FrozenTrial object
"""
return self._study.get_trials(deepcopy, states)
def ask(self, fixed_distributions: dict[str, BaseDistribution] | None = None) -> optuna.Trial:
"""Create a new trial from which hyperparameters can be suggested.
.. seealso::
See `Study.ask`_ for details.
.. _Study.ask: https://optuna.readthedocs.io/en/stable/reference/\
generated/optuna.study.Study.html#optuna.study.Study.ask
Args:
fixed_distributions:
A dictionary containing the parameter names and parameter's distributions. Each
parameter in this dictionary is automatically suggested for the returned trial,
even when the suggest method is not explicitly invoked by the user. If this
argument is set to :obj:`None`, no parameter is automatically suggested.
Returns:
A Trial object.
"""
return self._study.ask(fixed_distributions)
def add_trial(self, trial: FrozenTrial) -> None:
"""Add a trial to the study.
.. seealso::
See `Study.add_trials()`_ for details.
.. _Study.add_trials(): https://optuna.readthedocs.io/en/stable/reference/\
generated/optuna.study.Study.html#optuna.study.Study.add_trials
"""
self._study.add_trial(trial)
def add_trials(self, trials: Iterable[FrozenTrial]) -> None:
"""Add trials to the study.
.. seealso::
See `Study.add_trials()`_ for details.
.. _Study.add_trials(): https://optuna.readthedocs.io/en/stable/reference/\
generated/optuna.study.Study.html#optuna.study.Study.add_trials
"""
self._study.add_trials(trials)
def report_preference(
@@ -66,6 +185,14 @@ class PreferentialStudy:
better_trials: FrozenTrial | list[FrozenTrial],
worse_trials: FrozenTrial | list[FrozenTrial],
) -> None:
"""Report results of pairwise comparison.
Args:
better_trials:
Trials that are better than worse_trials.
worse_trials:
Trials that are worse than better_trials.
"""
if not isinstance(better_trials, list):
better_trials = [better_trials]
if not isinstance(worse_trials, list):
@@ -78,14 +205,43 @@ class PreferentialStudy:
)
def get_preferences(self, *, deepcopy: bool = True) -> list[tuple[FrozenTrial, FrozenTrial]]:
"""Return results of pairwise comparison.
Args:
deepcopy:
Flag to control whether to apply ``copy.deepcopy()`` to the trials.
Note that if you set the flag to :obj:`False`, you shouldn't mutate
any fields of the returned trial. Otherwise the internal state of
the study may corrupt and unexpected behavior may happen.
Returns:
A list of the pair of FrozenTrial objects. The left trial is better than the right one.
"""
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:
"""Set a user attribute to the study.
Args:
key: A key string of the attribute.
value: A value of the attribute. The value should be JSON serializable.
.. seealso::
See the `tutorial for user attributes <https://optuna.readthedocs.io/en/stable/\
tutorial/20_recipes/003_attributes.html>`_ on Optuna's documentation.
"""
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
@@ -99,18 +255,21 @@ class PreferentialStudy:
def get_best_trials(study_id: int, storage: optuna.storages.BaseStorage) -> list[FrozenTrial]:
ready_trials = [
t
for t in storage.get_all_trials(
study_id,
deepcopy=False,
states=(TrialState.COMPLETE, TrialState.RUNNING),
)
if t.system_attrs.get(_SYSTEM_ATTR_COMPARISON_READY) is True
]
preferences = get_preferences(study_id, storage)
worse_numbers = {worse for _, worse in preferences}
return [copy.deepcopy(t) for t in ready_trials if t.number not in worse_numbers]
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
if is_skipped_trial(t._trial_id, study_system_attrs):
continue
best_trials.append(copy.deepcopy(t))
return best_trials
def create_study(
@@ -120,6 +279,46 @@ def create_study(
study_name: str | None = None,
load_if_exists: bool = False,
) -> PreferentialStudy:
"""Like ``optuna.create_study()``, but for preferential optimization.
Example:
.. testcode::
import optuna
from optuna_dashboard.preferential import create_study
study = create_study()
trial = study.ask()
Args:
storage:
Database URL. If this argument is set to None, in-memory storage is used, and the
:class:`~optuna_dashboard.preferential.PreferentialStudy` will not be persistent.
sampler:
A sampler object that implements background algorithm for value suggestion.
If :obj:`None` is specified, `RandomSampler`_ is used. Please note that
most Optuna samplers does not work efficiently for preferential optimization.
.. _RandomSampler: https://optuna.readthedocs.io/en/stable/reference/\
samplers/generated/optuna.samplers.RandomSampler.html
study_name:
Study's name. If this argument is set to None, a unique name is generated
automatically.
load_if_exists:
Flag to control the behavior to handle a conflict of study names.
In the case where a study named ``study_name`` already exists in the ``storage``,
a :class:`~optuna.exceptions.DuplicatedStudyError` is raised if ``load_if_exists`` is
set to :obj:`False`.
Otherwise, the creation of the study is skipped, and the existing one is returned.
Returns:
A :class:`~optuna_dashboard.preferential.PreferentialStudy` object.
"""
try:
study = optuna.create_study(
storage=storage,
@@ -155,6 +354,52 @@ def load_study(
storage: str | optuna.storages.BaseStorage,
sampler: BaseSampler | None = None,
) -> PreferentialStudy:
"""Like ``optuna.load_study()``, but for preferential optimization.
Example:
.. testsetup::
import os
if os.path.exists("example.db"):
raise RuntimeError("'example.db' already exists. Please remove it.")
.. testcode::
import optuna
from optuna_dashboard.preferential import create_study
from optuna_dashboard.preferential import load_study
study = create_study(storage="sqlite:///example.db", study_name="my_study")
study.ask()
loaded_study = load_study(study_name="my_study", storage="sqlite:///example.db")
assert len(loaded_study.trials) == len(study.trials)
.. testcleanup::
os.remove("example.db")
Args:
study_name:
Study's name. Each study has a unique name as an identifier. If :obj:`None`, checks
whether the storage contains a single study, and if so loads that study.
``study_name`` is required if there are multiple studies in the storage.
storage:
Database URL such as ``sqlite:///example.db``. Please see also the documentation of
:func:`~optuna.study.create_study` for further details.
sampler:
A sampler object that implements background algorithm for value suggestion.
If :obj:`None` is specified, `RandomSampler`_ is used. Please note that
most Optuna samplers does not work efficiently for preferential optimization.
.. _RandomSampler: https://optuna.readthedocs.io/en/stable/reference/samplers/\
generated/optuna.samplers.RandomSampler.html
Returns:
A :class:`~optuna_dashboard.preferential.PreferentialStudy` object.
"""
study = optuna.load_study(
study_name=study_name, storage=storage, sampler=sampler or RandomSampler()
)
+20 -4
View File
@@ -1,14 +1,14 @@
from __future__ import annotations
from typing import Any
import uuid
from optuna.storages import BaseStorage
from optuna.trial import TrialState
from .._storage import get_study_summary
_SYSTEM_ATTR_PREFIX_PREFERENCE = "preference:values"
_SYSTEM_ATTR_PREFIX_SKIP_TRIAL = "preference:skip_trial:"
def report_preferences(
@@ -37,10 +37,26 @@ def get_preferences(
storage: BaseStorage,
) -> list[tuple[int, int]]:
preferences: list[tuple[int, int]] = []
summary = get_study_summary(storage, study_id)
system_attrs = getattr(summary, "system_attrs", {})
system_attrs = storage.get_study_system_attrs(study_id)
for k, v in system_attrs.items():
if not k.startswith(_SYSTEM_ATTR_PREFIX_PREFERENCE):
continue
preferences.extend(v) # type: ignore
return preferences
def report_skip(
study_id: int,
trial_id: int,
storage: BaseStorage,
) -> None:
storage.set_study_system_attr(
study_id=study_id,
key=_SYSTEM_ATTR_PREFIX_SKIP_TRIAL + str(trial_id),
value=True,
)
def is_skipped_trial(trial_id: int, study_system_attrs: dict[str, Any]) -> bool:
key = _SYSTEM_ATTR_PREFIX_SKIP_TRIAL + str(trial_id)
return key in study_system_attrs
@@ -1,6 +1 @@
from __future__ import annotations
from ._gp import PreferentialGPSampler
__all__ = ["PreferentialGPSampler"]
+12
View File
@@ -15,6 +15,7 @@ import {
getMetaInfoAPI,
deleteArtifactAPI,
reportPreferenceAPI,
skipPreferentialTrialAPI,
} from "./apiClient"
import {
graphVisibilityState,
@@ -598,6 +599,16 @@ export const actionCreator = () => {
})
}
const skipPreferentialTrial = (studyId: number, trialId: number) => {
skipPreferentialTrialAPI(studyId, trialId).catch((err) => {
const reason = err.response?.data.reason
enqueueSnackbar(`Failed to skip trial. Reason: ${reason}`, {
variant: "error",
})
console.log(err)
})
}
return {
updateAPIMeta,
updateStudyDetail,
@@ -618,6 +629,7 @@ export const actionCreator = () => {
makeTrialFail,
saveTrialUserAttrs,
updatePreference,
skipPreferentialTrial,
}
}
+13
View File
@@ -59,6 +59,7 @@ interface StudyDetailResponse {
name: string
datetime_start: string
directions: StudyDirection[]
user_attrs: Attribute[]
trials: TrialResponse[]
best_trials: TrialResponse[]
intersection_search_space: SearchSpaceItem[]
@@ -93,6 +94,7 @@ export const getStudyDetailAPI = (
name: res.data.name,
datetime_start: new Date(res.data.datetime_start),
directions: res.data.directions,
user_attrs: res.data.user_attrs,
trials: trials,
best_trials: best_trials,
union_search_space: res.data.union_search_space,
@@ -324,3 +326,14 @@ export const reportPreferenceAPI = (
return
})
}
export const skipPreferentialTrialAPI = (
studyId: number,
trialId: number
): Promise<void> => {
return axiosInstance
.post<void>(`/api/studies/${studyId}/${trialId}/skip`)
.then(() => {
return
})
}
@@ -12,6 +12,7 @@ import {
useTheme,
IconButton,
} from "@mui/material"
import Grid2 from "@mui/material/Unstable_Grid2"
import ChevronRightIcon from "@mui/icons-material/ChevronRight"
import Chip from "@mui/material/Chip"
import FormControlLabel from "@mui/material/FormControlLabel"
@@ -26,7 +27,7 @@ import HomeIcon from "@mui/icons-material/Home"
import { actionCreator } from "../action"
import { studySummariesState, studyDetailsState } from "../state"
import { AppDrawer } from "./AppDrawer"
import { GraphEdfMultiStudies } from "./GraphEdf"
import { GraphEdf } from "./GraphEdf"
import { GraphHistory } from "./GraphHistory"
import { useNavigate, useLocation } from "react-router-dom"
@@ -325,19 +326,21 @@ const StudiesGraph: FC<{ studies: StudySummary[] }> = ({ studies }) => {
</CardContent>
</Card>
) : null}
{showStudyDetails !== null &&
showStudyDetails.length > 0 &&
showStudyDetails.every((s) => s) ? (
<Card
sx={{
margin: theme.spacing(2),
}}
>
<CardContent>
<GraphEdfMultiStudies studies={showStudyDetails} />
</CardContent>
</Card>
) : null}
<Grid2 container spacing={2} sx={{ padding: theme.spacing(0, 2) }}>
{showStudyDetails !== null &&
showStudyDetails.length > 0 &&
showStudyDetails.every((s) => s)
? showStudyDetails[0].directions.map((d, i) => (
<Grid2 xs={6} key={i}>
<Card>
<CardContent>
<GraphEdf studies={showStudyDetails} objectiveId={i} />
</CardContent>
</Card>
</Grid2>
))
: null}
</Grid2>
</Box>
)
}
+15 -147
View File
@@ -1,25 +1,9 @@
import * as plotly from "plotly.js-dist-min"
import React, { FC, useEffect, useMemo } from "react"
import {
Grid,
FormControl,
FormLabel,
MenuItem,
Select,
Typography,
SelectChangeEvent,
useTheme,
Box,
} from "@mui/material"
import { Typography, useTheme, Box } from "@mui/material"
import { plotlyDarkTemplate } from "./PlotlyDarkMode"
import {
Target,
useFilteredTrials,
useFilteredTrialsFromStudies,
useObjectiveTargets,
} from "../trialFilter"
import { Target, useFilteredTrialsFromStudies } from "../trialFilter"
const plotDomId = "graph-edf"
const getPlotDomId = (objectiveId: number) => `graph-edf-${objectiveId}`
interface EdfPlotInfo {
@@ -28,44 +12,16 @@ interface EdfPlotInfo {
}
export const GraphEdf: FC<{
study: StudyDetail | null
studies: StudyDetail[]
objectiveId: number
}> = ({ study, objectiveId }) => {
}> = ({ studies, objectiveId }) => {
const theme = useTheme()
const domId = getPlotDomId(objectiveId)
const target = useMemo<Target>(
() => new Target("objective", objectiveId),
[objectiveId]
)
const trials = useFilteredTrials(study, [target], false)
useEffect(() => {
if (study !== null) {
plotEdf(trials, target, domId, theme.palette.mode)
}
}, [trials, target, domId, theme.palette.mode])
return (
<Box>
<Typography
variant="h6"
sx={{ margin: "1em 0", fontWeight: theme.typography.fontWeightBold }}
>
{`EDF for ${target.toLabel(study?.objective_names)}`}
</Typography>
<Box id={domId} sx={{ height: "450px" }} />
</Box>
)
}
export const GraphEdfMultiStudies: FC<{
studies: StudyDetail[]
}> = ({ studies }) => {
const theme = useTheme()
const [targets, selected, setTarget] = useObjectiveTargets(
studies.length !== 0 ? studies[0] : null
)
const trials = useFilteredTrialsFromStudies(studies, [selected], false)
const trials = useFilteredTrialsFromStudies(studies, [target], false)
const edfPlotInfos = studies.map((study, index) => {
const e: EdfPlotInfo = {
study_name: study?.name,
@@ -74,112 +30,24 @@ export const GraphEdfMultiStudies: FC<{
return e
})
const handleObjectiveChange = (event: SelectChangeEvent<string>) => {
setTarget(event.target.value)
}
useEffect(() => {
plotEdfMultiStudies(edfPlotInfos, selected, plotDomId, theme.palette.mode)
}, [studies, selected, theme.palette.mode])
plotEdf(edfPlotInfos, target, domId, theme.palette.mode)
}, [studies, target, theme.palette.mode])
return (
<Grid container direction="row">
<Grid
item
xs={3}
container
direction="column"
sx={{ paddingRight: theme.spacing(2) }}
<Box>
<Typography
variant="h6"
sx={{ margin: "1em 0", fontWeight: theme.typography.fontWeightBold }}
>
<Typography
variant="h6"
sx={{ margin: "1em 0", fontWeight: theme.typography.fontWeightBold }}
>
EDF
</Typography>
{studies.length > 0 && studies[0].directions.length !== 1 ? (
<FormControl component="fieldset">
<FormLabel component="legend">Objective:</FormLabel>
<Select
value={selected.identifier()}
onChange={handleObjectiveChange}
>
{targets.map((target, i) => (
<MenuItem value={target.identifier()} key={i}>
{target.toLabel(studies[0].objective_names)}
</MenuItem>
))}
</Select>
</FormControl>
) : null}
</Grid>
<Grid item xs={9}>
<Box id={plotDomId} sx={{ height: "450px" }} />
</Grid>
</Grid>
{`EDF for ${target.toLabel(studies[0].objective_names)}`}
</Typography>
<Box id={domId} sx={{ height: "450px" }} />
</Box>
)
}
const plotEdf = (
trials: Trial[],
target: Target,
domId: string,
mode: string
) => {
if (document.getElementById(domId) === null) {
return
}
if (trials.length === 0) {
plotly.react(domId, [], {
template: mode === "dark" ? plotlyDarkTemplate : {},
})
return
}
const target_name = "Objective Value"
const layout: Partial<plotly.Layout> = {
xaxis: {
title: target_name,
},
yaxis: {
title: "Cumulative Probability",
},
margin: {
l: 50,
t: 0,
r: 50,
b: 50,
},
uirevision: "true",
template: mode === "dark" ? plotlyDarkTemplate : {},
}
const values = trials.map((t) => target.getTargetValue(t) as number)
const numValues = values.length
const minX = Math.min(...values)
const maxX = Math.max(...values)
const numStep = 100
const _step = (maxX - minX) / (numStep - 1)
const xValues = []
const yValues = []
for (let i = 0; i < numStep; i++) {
const boundary_right = minX + _step * i
xValues.push(boundary_right)
yValues.push(values.filter((v) => v <= boundary_right).length / numValues)
}
const plotData: Partial<plotly.PlotData>[] = [
{
type: "scatter",
x: xValues,
y: yValues,
},
]
plotly.react(domId, plotData, layout)
}
const plotEdfMultiStudies = (
edfPlotInfos: EdfPlotInfo[],
target: Target,
domId: string,
+1 -1
View File
@@ -163,7 +163,7 @@ const useConfirmCloseDialog = (
return [openDialog, renderDialog]
}
const MarkdownRenderer: FC<{ body: string }> = ({ body }) => (
export const MarkdownRenderer: FC<{ body: string }> = ({ body }) => (
<ReactMarkdown
children={body}
remarkPlugins={[remarkGfm, remarkMath]}
@@ -1,8 +1,23 @@
import React, { FC, useState } from "react"
import { Typography, Box, Button, useTheme } from "@mui/material"
import {
Typography,
Box,
useTheme,
Card,
CardContent,
CardActions,
CardActionArea,
} from "@mui/material"
import ClearIcon from "@mui/icons-material/Clear"
import IconButton from "@mui/material/IconButton"
import OpenInFullIcon from "@mui/icons-material/OpenInFull"
import ReplayIcon from "@mui/icons-material/Replay"
import Modal from "@mui/material/Modal"
import { red } from "@mui/material/colors"
import { TrialNote } from "./Note"
import { actionCreator } from "../action"
import { TrialListDetail } from "./TrialList"
import { MarkdownRenderer } from "./Note"
const PreferentialTrial: FC<{
trial?: Trial
@@ -12,41 +27,149 @@ const PreferentialTrial: FC<{
const theme = useTheme()
const action = actionCreator()
const trialWidth = 500
const trialHeight = 300
const [detailShown, setDetailShown] = useState(false)
if (trial == undefined) {
return <Box width={trialWidth}></Box>
return (
<Box
sx={{
width: trialWidth,
minHeight: trialHeight,
margin: theme.spacing(2),
}}
/>
)
}
return (
<Box sx={{ width: trialWidth, padding: theme.spacing(2, 2, 0, 2) }}>
<Typography
variant="h4"
sx={{
marginBottom: theme.spacing(2),
fontWeight: theme.typography.fontWeightBold,
}}
>
Trial {trial.number} (trial_id={trial.trial_id})
</Typography>
<Button
variant="outlined"
onClick={() => {
hideTrial()
const best_trials = studyDetail.best_trials
.map((t) => t.number)
.filter((t) => t !== trial.number)
action.updatePreference(trial.study_id, best_trials, [trial.number])
}}
>
Worst
</Button>
<TrialNote
studyId={trial.study_id}
trialId={trial.trial_id}
latestNote={trial.note}
cardSx={{ marginBottom: theme.spacing(2) }}
/>
</Box>
<Card
sx={{
width: trialWidth,
minHeight: trialHeight,
margin: theme.spacing(2),
padding: 0,
}}
>
<CardActions>
<Typography variant="h5">Trial {trial.number}</Typography>
<IconButton
sx={{
marginLeft: "auto",
}}
onClick={() => {
hideTrial()
action.skipPreferentialTrial(trial.study_id, trial.trial_id)
}}
aria-label="skip trial"
>
<ReplayIcon />
</IconButton>
<IconButton
sx={{
marginLeft: "auto",
}}
onClick={() => setDetailShown(true)}
aria-label="show detail"
>
<OpenInFullIcon />
</IconButton>
</CardActions>
<CardActionArea>
<CardContent
aria-label="trial-button"
onClick={() => {
hideTrial()
const best_trials = studyDetail.best_trials
.map((t) => t.number)
.filter((t) => t !== trial.number)
action.updatePreference(trial.study_id, best_trials, [trial.number])
}}
sx={{
padding: 0,
position: "relative",
overflow: "hidden",
"::before": {
content: '""',
position: "absolute",
top: 0,
left: 0,
width: "100%",
height: "100%",
backgroundColor:
theme.palette.mode === "dark" ? "white" : "black",
opacity: 0,
zIndex: 1,
transition: "opacity 0.3s ease-out",
},
":hover::before": {
opacity: 0.2,
},
}}
>
<Box
sx={{
padding: theme.spacing(2),
}}
>
<MarkdownRenderer body={trial.note.body} />
</Box>
<ClearIcon
sx={{
position: "absolute",
width: "100%",
height: "100%",
top: 0,
left: 0,
color: red[600],
opacity: 0,
transition: "opacity 0.3s ease-out",
zIndex: 1,
":hover": {
opacity: 0.3,
filter:
theme.palette.mode === "dark"
? "brightness(1.1)"
: "brightness(1.7)",
},
}}
/>
</CardContent>
</CardActionArea>
<Modal open={detailShown} onClose={() => setDetailShown(false)}>
<Box
sx={{
position: "absolute",
top: 0,
left: 0,
right: 0,
bottom: 0,
width: "80%",
maxHeight: "90%",
margin: "auto",
overflow: "hidden",
backgroundColor: theme.palette.mode === "dark" ? "black" : "white",
borderRadius: theme.spacing(3),
}}
>
<Box
sx={{
width: "100%",
height: "100%",
overflow: "auto",
}}
>
<TrialListDetail
trial={trial}
isBestTrial={() => true}
directions={[]}
objectiveNames={[]}
/>
</Box>
</Box>
</Modal>
</Card>
)
}
@@ -61,6 +184,7 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({
if (studyDetail === null || !studyDetail.is_preferential) {
return null
}
const theme = useTheme()
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),
@@ -104,17 +228,28 @@ export const PreferentialTrials: FC<{ studyDetail: StudyDetail | null }> = ({
}
return (
<Box sx={{ display: "flex", flexDirection: "row", flexWrap: "wrap" }}>
{displayTrials.numbers.map((t, index) => (
<PreferentialTrial
key={index}
trial={studyDetail.best_trials.find((trial) => trial.number === t)}
studyDetail={studyDetail}
hideTrial={() => {
hideTrial(t)
}}
/>
))}
<Box padding={theme.spacing(2)}>
<Typography
variant="h4"
sx={{
marginBottom: theme.spacing(2),
fontWeight: theme.typography.fontWeightBold,
}}
>
Which trial is the worst?
</Typography>
<Box sx={{ display: "flex", flexDirection: "row", flexWrap: "wrap" }}>
{displayTrials.numbers.map((t, index) => (
<PreferentialTrial
key={index}
trial={studyDetail.best_trials.find((trial) => trial.number === t)}
studyDetail={studyDetail}
hideTrial={() => {
hideTrial(t)
}}
/>
))}
</Box>
</Box>
)
}
@@ -127,7 +127,7 @@ export const StudyDetail: FC<{
<Grid2 xs={6} key={i}>
<Card>
<CardContent>
<GraphEdf study={studyDetail} objectiveId={i} />
<GraphEdf studies={[studyDetail]} objectiveId={i} />
</CardContent>
</Card>
</Grid2>
@@ -39,7 +39,7 @@ export const StudyHistory: FC<{ studyId: number }> = ({ studyId }) => {
setIncludePruned(!includePruned)
}
const userAttrs = studySummary?.user_attrs || []
const userAttrs = studySummary?.user_attrs || studyDetail?.user_attrs || []
const userAttrColumns: DataGridColumn<Attribute>[] = [
{ field: "key", label: "Key", sortable: true },
{ field: "value", label: "Value", sortable: true },
+44 -42
View File
@@ -143,7 +143,7 @@ const useIsBestTrial = (
}, [studyDetail])
}
const TrialListDetail: FC<{
export const TrialListDetail: FC<{
trial: Trial
isBestTrial: (trialId: number) => boolean
directions: StudyDirection[]
@@ -710,55 +710,57 @@ const TrialArtifact: FC<{ trial: Trial }> = ({ trial }) => {
)
}
})}
<Card
sx={{
marginBottom: theme.spacing(2),
width: width,
minHeight: height,
margin: theme.spacing(0, 1, 1, 0),
border: dragOver
? `3px dashed ${
theme.palette.mode === "dark" ? "white" : "black"
}`
: `1px solid ${theme.palette.divider}`,
}}
onDragOver={handleDragOver}
onDragLeave={handleDragLeave}
onDrop={handleDrop}
>
<CardActionArea
onClick={handleClick}
{trial.state === "Running" || trial.state === "Waiting" ? (
<Card
sx={{
height: "100%",
marginBottom: theme.spacing(2),
width: width,
minHeight: height,
margin: theme.spacing(0, 1, 1, 0),
border: dragOver
? `3px dashed ${
theme.palette.mode === "dark" ? "white" : "black"
}`
: `1px solid ${theme.palette.divider}`,
}}
onDragOver={handleDragOver}
onDragLeave={handleDragLeave}
onDrop={handleDrop}
>
<CardContent
<CardActionArea
onClick={handleClick}
sx={{
display: "flex",
height: "100%",
flexDirection: "column",
justifyContent: "center",
alignItems: "center",
}}
>
<UploadFileIcon
sx={{ fontSize: 80, marginBottom: theme.spacing(2) }}
/>
<input
type="file"
ref={inputRef}
onChange={handleOnChange}
style={{ display: "none" }}
/>
<Typography>Upload a New File</Typography>
<Typography
sx={{ textAlign: "center", color: theme.palette.grey.A400 }}
<CardContent
sx={{
display: "flex",
height: "100%",
flexDirection: "column",
justifyContent: "center",
alignItems: "center",
}}
>
Drag your file here or click to browse.
</Typography>
</CardContent>
</CardActionArea>
</Card>
<UploadFileIcon
sx={{ fontSize: 80, marginBottom: theme.spacing(2) }}
/>
<input
type="file"
ref={inputRef}
onChange={handleOnChange}
style={{ display: "none" }}
/>
<Typography>Upload a New File</Typography>
<Typography
sx={{ textAlign: "center", color: theme.palette.grey.A400 }}
>
Drag your file here or click to browse.
</Typography>
</CardContent>
</CardActionArea>
</Card>
) : null}
</Box>
{renderDeleteArtifactDialog()}
</>
@@ -157,7 +157,7 @@ export const TrialTable: FC<{
if (firstVal === secondVal) {
return 0
} else if (firstVal && secondVal) {
return firstVal < secondVal ? 1 : -1
return Number(firstVal) < Number(secondVal) ? 1 : -1
} else if (firstVal) {
return -1
} else {
+1
View File
@@ -185,6 +185,7 @@ type StudyDetail = {
id: number
name: string
directions: StudyDirection[]
user_attrs: Attribute[]
datetime_start: Date
best_trials: Trial[]
trials: Trial[]
+82
View File
@@ -0,0 +1,82 @@
from unittest.mock import MagicMock
from optuna.storages import BaseStorage
from optuna_dashboard.artifact import _backend
import pytest
def test_get_artifact_path() -> None:
study = MagicMock(_study_id=0)
trial = MagicMock(_trial_id=0, study=study)
assert _backend.get_artifact_path(trial=trial, artifact_id="id0") == "/artifacts/0/0/id0"
def test_artifact_prefix() -> None:
actual = _backend._artifact_prefix(trial_id=0)
assert actual == "dashboard:artifacts:0:"
@pytest.fixture()
def init_storage_with_artifact_meta() -> BaseStorage:
from optuna import create_study
from optuna.storages import InMemoryStorage
storage = InMemoryStorage()
study = create_study(storage=storage)
study_system_attrs = {
"dashboard:artifacts:0:id0": '{"artifact_id": "id0", "filename": "foo.txt"}',
"dashboard:artifacts:0:id1": '{"artifact_id": "id1", "filename": "bar.txt"}',
"baz": "baz",
}
for key, value in study_system_attrs.items():
study.set_system_attr(key, value)
trial_system_attrs = {
"artifacts:id2": '{"artifact_id": "id2", "filename": "baz.txt"}',
"artifacts:id3": '{"artifact_id": "id3", "filename": "qux.txt"}',
}
for key, value in trial_system_attrs.items():
trial = study.ask()
trial.set_system_attr(key, value)
study.tell(trial, 0.0)
return storage
def test_get_artifact_meta(init_storage_with_artifact_meta: MagicMock) -> None:
storage = init_storage_with_artifact_meta
actual = _backend.get_artifact_meta(storage, study_id=0, trial_id=0, artifact_id="id0")
assert actual == {"artifact_id": "id0", "filename": "foo.txt"}
actual = _backend.get_artifact_meta(storage, study_id=0, trial_id=1, artifact_id="id3")
assert actual == {"artifact_id": "id3", "filename": "qux.txt"}
actual = _backend.get_artifact_meta(storage, study_id=0, trial_id=0, artifact_id="id4")
assert actual is None
def test_delete_all_artifacts(init_storage_with_artifact_meta: MagicMock) -> None:
backend = MagicMock()
storage = init_storage_with_artifact_meta
_backend.delete_all_artifacts(backend, storage, study_id=0)
assert backend.remove.call_args_list == [
(("id0",),),
(("id1",),),
(("id2",),),
(("id3",),),
]
def test_list_trial_artifacts(init_storage_with_artifact_meta: MagicMock) -> None:
storage = init_storage_with_artifact_meta
trial = MagicMock(_trial_id=0, system_attrs=storage.get_trial_system_attrs(0))
actual = _backend.list_trial_artifacts(storage.get_study_system_attrs(0), trial)
assert actual == [
{"artifact_id": "id0", "filename": "foo.txt"},
{"artifact_id": "id1", "filename": "bar.txt"},
{"artifact_id": "id2", "filename": "baz.txt"},
]
+24
View File
@@ -151,6 +151,30 @@ class APITestCase(TestCase):
assert better.number == 2
assert worse.number == 1
def test_skip_trial(self) -> None:
storage = optuna.storages.InMemoryStorage()
study = create_study(storage=storage)
trials: list[optuna.Trial] = []
for _ in range(3):
trial = study.ask()
study.mark_comparison_ready(trial)
trials.append(trial)
app = create_app(storage)
study_id = study._study._study_id
status, _, _ = send_request(
app,
f"/api/studies/{study_id}/{trials[1]._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
def test_create_study(self) -> None:
for name, directions, expected_status in [
("single-objective success", ["minimize"], 201),
@@ -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)
}
+26 -7
View File
@@ -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>
)
+60 -30
View File
@@ -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}
+272 -144
View File
@@ -2,6 +2,14 @@
import sqlite3InitModule from "@sqlite.org/sqlite-wasm"
import { SetterOrUpdater } from "recoil"
type SQLite3DB = {
exec(options: {
sql: string
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (...args: any[]) => void
}): void
}
export const loadStorage = (
arrayBuffer: ArrayBuffer,
setter: SetterOrUpdater<Study[]>
@@ -30,155 +38,275 @@ export const loadStorage = (
)
db.checkRc(rc)
try {
// Check version_info table
let supported = true
db.exec({
sql: "SELECT schema_version FROM version_info LIMIT 1",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
if (vals[0] != 12) {
supported = false
}
},
})
if (!supported) {
if (!isSupportedSchema(db)) {
return
}
// Get studies
const studies: Study[] = []
db.exec({
sql:
"SELECT s.study_id, s.study_name, sd.direction, sd.objective" +
" FROM studies AS s INNER JOIN study_directions AS sd" +
" ON s.study_id = sd.study_id ORDER BY sd.study_direction_id",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const study_id = vals[0]
const study_name = vals[1]
const direction: StudyDirection = vals[2].toLowerCase()
const objective = vals[3]
let index = 0
if (objective === 0) {
studies.push({
study_id: study_id,
study_name: study_name,
directions: [direction],
union_search_space: [],
intersection_search_space: [],
user_attrs: [],
system_attrs: [],
trials: [],
})
} else {
index = studies.findIndex((s) => s.study_id === study_id)
studies[index].directions.push(direction)
}
},
})
studies.forEach((s) => {
db.exec({
sql:
"SELECT t.trial_id, t.number, t.study_id, t.state, t.datetime_start, t.datetime_complete," +
" tv.objective, tv.value, tv.value_type" +
" FROM trials AS t LEFT JOIN trial_values AS tv ON tv.trial_id = t.trial_id" +
` WHERE t.study_id = ${s.study_id}` +
" ORDER BY t.number",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const state: TrialState =
vals[3] === "COMPLETE"
? "Complete"
: vals[3] === "PRUNED"
? "Pruned"
: vals[3] === "RUNNING"
? "Running"
: vals[3] === "WAITING"
? "Waiting"
: "Fail"
const trial: Trial = {
trial_id: vals[0],
number: vals[1],
study_id: vals[2],
state: state,
params: [],
intermediate_values: [],
user_attrs: [],
system_attrs: [],
}
s.trials.push(trial)
},
})
const union_search_space: SearchSpaceItem[] = []
let intersection_search_space: Set<SearchSpaceItem> = new Set()
s.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]
params.push({
name: param_name,
param_internal_value: vals[1],
})
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) => {
intersection_search_space.add({
name: s,
})
})
} else {
intersection_search_space = new Set(
Array.from(intersection_search_space).filter((s) =>
param_names.has(s.name)
)
)
}
trial.params = params
const values: TrialValueNumber[] = []
db.exec({
sql:
"SELECT value, value_type" +
` FROM trial_values WHERE trial_id = ${trial.trial_id}` +
" ORDER BY objective",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
values.push(
vals[1] === "INF_NEG"
? "-inf"
: vals[1] === "INF_POS"
? "+inf"
: vals[0]
)
},
})
if (s.directions.length === values.length) {
trial.values = values
}
})
s.union_search_space = union_search_space
s.intersection_search_space = Array.from(intersection_search_space)
})
const studies = getStudies(db)
setter((prev) => [...prev, ...studies])
} finally {
db.close()
}
})
}
const isSupportedSchema = (db: SQLite3DB): boolean => {
let supported = true
db.exec({
sql: "SELECT schema_version FROM version_info LIMIT 1",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
if (vals[0] != 12) {
supported = false
}
},
})
return supported
}
const getStudies = (db: SQLite3DB): Study[] => {
const studies: Study[] = []
db.exec({
sql:
"SELECT s.study_id, s.study_name, sd.direction, sd.objective" +
" FROM studies AS s INNER JOIN study_directions AS sd" +
" ON s.study_id = sd.study_id ORDER BY sd.study_direction_id",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const studyId = vals[0]
const studyName = vals[1]
const direction: StudyDirection =
vals[2] === "MINIMIZE" ? "minimize" : "maximize"
const objective = vals[3]
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 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 })
}
})
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) => {
intersection_search_space.add({
name: s,
})
})
} else {
intersection_search_space = new Set(
Array.from(intersection_search_space).filter((s) =>
param_names.has(s.name)
)
)
}
trial.params = params
trial.user_attrs = userAttrs
})
if (objective === 0) {
studies.push({
study_id: studyId,
study_name: studyName,
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
}
const index = studies.findIndex((s) => s.study_id === studyId)
studies[index].directions.push(direction)
},
})
return studies
}
const getTrials = (db: SQLite3DB, studyId: number): Trial[] => {
const trials: Trial[] = []
db.exec({
sql:
"SELECT trial_id, number, state, datetime_start, datetime_complete FROM trials" +
` WHERE study_id = ${studyId} ORDER BY number`,
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
const trialId = vals[0]
const state: TrialState =
vals[2] === "COMPLETE"
? "Complete"
: vals[2] === "PRUNED"
? "Pruned"
: vals[2] === "RUNNING"
? "Running"
: vals[2] === "WAITING"
? "Waiting"
: "Fail"
const trial: Trial = {
trial_id: trialId,
number: vals[1],
study_id: studyId,
state: state,
values: getTrialValues(db, trialId),
intermediate_values: getTrialIntermediateValues(db, trialId),
params: [], // Set this column later
user_attrs: [], // Set this column later
datetime_start: vals[3],
datetime_complete: vals[4],
}
trials.push(trial)
},
})
return trials
}
const getTrialValues = (db: SQLite3DB, trialId: number): TrialValueNumber[] => {
const values: TrialValueNumber[] = []
db.exec({
sql:
"SELECT value, value_type" +
` FROM trial_values WHERE trial_id = ${trialId}` +
" ORDER BY objective",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
values.push(
vals[1] === "INF_NEG"
? "-inf"
: vals[1] === "INF_POS"
? "+inf"
: vals[0]
)
},
})
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
): TrialIntermediateValue[] => {
const values: TrialIntermediateValue[] = []
db.exec({
sql:
"SELECT step, intermediate_value, intermediate_value_type" +
` FROM trial_intermediate_values WHERE trial_id = ${trialId}` +
" ORDER BY step",
// eslint-disable-next-line @typescript-eslint/no-explicit-any
callback: (vals: any[]) => {
values.push({
step: vals[0],
value:
vals[2] === "INF_NEG"
? "-inf"
: vals[2] === "INF_POS"
? "+inf"
: vals[1],
})
},
})
return values
}
+15 -4
View File
@@ -27,6 +27,11 @@ type CategoricalDistribution = {
choices: { pytype: string; value: string }[]
}
type TrialIntermediateValue = {
step: number
value: TrialIntermediateValueNumber
}
type Distribution =
| FloatDistribution
| IntDistribution
@@ -37,14 +42,18 @@ type Attribute = {
value: string
}
type AttributeSpec = {
key: string
sortable: boolean
}
type Study = {
study_id: number
study_name: string
directions: StudyDirection[]
user_attrs: Attribute[]
union_search_space: SearchSpaceItem[]
intersection_search_space: SearchSpaceItem[]
system_attrs: Attribute[]
union_user_attrs: AttributeSpec[]
datetime_start?: Date
trials: Trial[]
}
@@ -57,15 +66,17 @@ type Trial = {
values?: TrialValueNumber[]
params: TrialParam[]
intermediate_values: TrialIntermediateValue[]
user_attrs: Attribute[]
datetime_start?: Date
datetime_complete?: Date
user_attrs: Attribute[]
system_attrs: Attribute[]
}
type TrialParam = {
name: string
param_internal_value: number
param_external_value: string
param_external_type: string
distribution: Distribution
}
type SearchSpaceItem = {
+7 -7
View File
@@ -1,17 +1,17 @@
# optuna-dashboard README
# Optuna Dashboard for VS Code
## Features
The Optuna Dashboard extension lets you open Optuna's SQLite3 database in Optuna Dashboard, allowing users to see optimization histories in graphs and tables.
VSCode Extension that launches Optuna Dashboard (Wasm ver.).
![vscode-extension](https://github.com/optuna/optuna-dashboard/raw/main/docs/_static/vscode-extension.png)
## Usage
Please right-click on the SQLite3 file (`*.db` or `*.sqlite3`) in the file explorer and select the "Open in Optuna Dashboard" command from the dropdown menu.
## Extension Settings
Nothing to configure.
## Known Issues
* ...
## Release Notes
### 0.0.1
+8038 -8038
View File
File diff suppressed because it is too large. Load diff
+1
View File
@@ -2,6 +2,7 @@
"name": "optuna-dashboard",
"displayName": "Optuna Dashboard",
"description": "Web Dashboard for Optuna",
"publisher": "Optuna",
"version": "0.0.1",
"license": "MIT",
"icon": "images/optuna-logo.png",
+1 -1
View File
@@ -8,7 +8,7 @@ export function activate(context: vscode.ExtensionContext) {
let disposable = vscode.commands.registerCommand(
"optuna-dashboard.openOptunaDashboard",
async (fileUri: vscode.Uri) => {
// In VSCode, the path separator of fileUri is always '/'
// In VS Code, the path separator of fileUri is always '/'
// even when using Windows.
const title = fileUri.path.split("/").pop() || "Optuna Dashboard"
const panel = vscode.window.createWebviewPanel(