diff --git a/docs/api.rst b/docs/api.rst index 1d3ec40a..53b59d24 100644 --- a/docs/api.rst +++ b/docs/api.rst @@ -23,6 +23,7 @@ Human-in-the-loop :nosignatures: optuna_dashboard.register_objective_form_widgets + optuna_dashboard.dict_to_form_widget optuna_dashboard.ChoiceWidget optuna_dashboard.SliderWidget optuna_dashboard.TextInputWidget diff --git a/optuna_dashboard/__init__.py b/optuna_dashboard/__init__.py index e273992b..b1ea71c4 100644 --- a/optuna_dashboard/__init__.py +++ b/optuna_dashboard/__init__.py @@ -1,6 +1,7 @@ from ._app import run_server # noqa from ._app import wsgi # noqa from ._form_widget import ChoiceWidget # noqa +from ._form_widget import dict_to_form_widget # noqa from ._form_widget import ObjectiveChoiceWidget # noqa from ._form_widget import ObjectiveSliderWidget # noqa from ._form_widget import ObjectiveTextInputWidget # noqa diff --git a/optuna_dashboard/_form_widget.py b/optuna_dashboard/_form_widget.py index 7b076ec1..b3e76a48 100644 --- a/optuna_dashboard/_form_widget.py +++ b/optuna_dashboard/_form_widget.py @@ -76,6 +76,16 @@ class ChoiceWidget: "user_attr_key": self.user_attr_key, } + @classmethod + def _from_dict(cls, d: dict[str, Any]) -> ChoiceWidget: + assert d.get("type") == "choice" + return cls( + description=d.get("description"), + choices=d["choices"], + values=d["values"], + user_attr_key=d.get("user_attr_key"), + ) + @dataclass class SliderWidget: @@ -100,6 +110,21 @@ class SliderWidget: "user_attr_key": self.user_attr_key, } + @classmethod + def _from_dict(cls, d: dict[str, Any]) -> SliderWidget: + assert d.get("type") == "slider" + labels = d.get("labels") + if labels is not None: + labels = [(l["value"], l["label"]) for l in labels] + return cls( + description=d.get("description"), + min=d["min"], + max=d["max"], + step=d.get("step"), + labels=labels, + user_attr_key=d.get("user_attr_key"), + ) + @dataclass class TextInputWidget: @@ -113,10 +138,19 @@ class TextInputWidget: "user_attr_key": self.user_attr_key, } + @classmethod + def _from_dict(cls, d: dict[str, Any]) -> TextInputWidget: + assert d.get("type") == "text" + return cls( + description=d.get("description"), + user_attr_key=d.get("user_attr_key"), + ) + @dataclass class ObjectiveUserAttrRef: key: str + # TODO(c-bata): Remove this attribute user_attr_key: Optional[str] = None def to_dict(self) -> UserAttrRefJSON: @@ -126,6 +160,13 @@ class ObjectiveUserAttrRef: "user_attr_key": self.user_attr_key, } + @classmethod + def _from_dict(cls, d: dict[str, Any]) -> ObjectiveUserAttrRef: + assert d.get("type") == "user_attr" + return cls( + key=d["key"], + ) + ObjectiveFormWidget = Union[ChoiceWidget, SliderWidget, TextInputWidget, ObjectiveUserAttrRef] # For backward compatibility. @@ -135,6 +176,27 @@ ObjectiveTextInputWidget = TextInputWidget FORM_WIDGETS_KEY = "dashboard:form_widgets:v2" +def dict_to_form_widget(d: dict[str, Any]) -> ObjectiveFormWidget: + """Restore form widget objects from the dictionary. + + Args: + d: A dictionary object. + + Returns: + object: an instance of the restored form widget class. + """ + widget_type = d.get("type", None) + if widget_type == "choice": + return ChoiceWidget._from_dict(d) + elif widget_type == "slider": + return SliderWidget._from_dict(d) + elif widget_type == "text": + return TextInputWidget._from_dict(d) + elif widget_type == "user_attr": + return ObjectiveUserAttrRef._from_dict(d) + raise ValueError("Unexpected widget type") + + def register_objective_form_widgets( study: optuna.Study, widgets: list[ObjectiveFormWidget] ) -> None: diff --git a/python_tests/test_form_widget.py b/python_tests/test_form_widget.py new file mode 100644 index 00000000..68b4f19d --- /dev/null +++ b/python_tests/test_form_widget.py @@ -0,0 +1,40 @@ +from __future__ import annotations + +from unittest import TestCase + +from optuna_dashboard import ChoiceWidget +from optuna_dashboard import dict_to_form_widget +from optuna_dashboard import ObjectiveUserAttrRef +from optuna_dashboard import SliderWidget +from optuna_dashboard import TextInputWidget + + +class FormWidgetsTestCase(TestCase): + def test_widget_to_dict_from_dict(self) -> None: + widgets = [ + ChoiceWidget(choices=["Good", "Bad"], values=[1, -1]), + ChoiceWidget( + choices=["Good", "Bad"], + values=[1, -1], + description="description", + user_attr_key="key", + ), + SliderWidget(min=1, max=5), + SliderWidget( + min=1, + max=5, + step=1, + labels=[(1, "Bad"), (5, "Good")], + description="description", + user_attr_key="key", + ), + TextInputWidget(), + TextInputWidget(description="description", user_attr_key="key"), + ObjectiveUserAttrRef(key="key"), + ] + + for i, widget in enumerate(widgets): + with self.subTest(f"{widget.__class__}-{i}"): + d = widget.to_dict() + restored = dict_to_form_widget(d) + assert widget == restored