diff --git a/optuna_dashboard/__init__.py b/optuna_dashboard/__init__.py index 7b273ed1..cf75f553 100644 --- a/optuna_dashboard/__init__.py +++ b/optuna_dashboard/__init__.py @@ -7,6 +7,7 @@ from ._objective_form_widget import ObjectiveSliderWidget # noqa from ._objective_form_widget import ObjectiveTextInputWidget # noqa from ._objective_form_widget import ObjectiveUserAttrRef # noqa from ._objective_form_widget import register_objective_form_widgets # noqa +from ._objective_form_widget import register_user_attr_form_widgets # noqa __version__ = "0.9.0b6" diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index c521775a..723896cc 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -34,6 +34,7 @@ from ._bottle_util import BottleViewReturn from ._bottle_util import json_api_view from ._cached_extra_study_property import get_cached_extra_study_property from ._importance import get_param_importance_from_trials_cache +from ._objective_form_widget import SYSTEM_ATTR_OUTPUT_TYPE_KEY from ._pareto_front import get_pareto_front_trials from ._serializer import serialize_study_detail from ._serializer import serialize_study_summary @@ -357,6 +358,7 @@ def create_app( union_user_attrs, has_intermediate_values, ) = get_cached_extra_study_property(study_id, trials) + form_widgets_output_type = storage.get_study_system_attrs(study_id).get(SYSTEM_ATTR_OUTPUT_TYPE_KEY) return serialize_study_detail( summary, best_trials, @@ -365,6 +367,7 @@ def create_app( union, union_user_attrs, has_intermediate_values, + form_widgets_output_type, ) @app.get("/api/studies//param_importances") diff --git a/optuna_dashboard/_objective_form_widget.py b/optuna_dashboard/_objective_form_widget.py index 3f19631e..ea3d737c 100644 --- a/optuna_dashboard/_objective_form_widget.py +++ b/optuna_dashboard/_objective_form_widget.py @@ -111,6 +111,7 @@ ObjectiveFormWidget = Union[ ObjectiveChoiceWidget, ObjectiveSliderWidget, ObjectiveTextInputWidget, ObjectiveUserAttrRef ] SYSTEM_ATTR_KEY = "dashboard:objective_form_widgets:v1" +SYSTEM_ATTR_OUTPUT_TYPE_KEY = "dashboard:form_widgets_output_type:v1" def register_objective_form_widgets( @@ -120,6 +121,15 @@ def register_objective_form_widgets( raise ValueError("The length of actions must be the same with the number of objectives.") widget_dicts = [w.to_dict() for w in widgets] study._storage.set_study_system_attr(study._study_id, SYSTEM_ATTR_KEY, widget_dicts) + study._storage.set_study_system_attr(study._study_id, SYSTEM_ATTR_OUTPUT_TYPE_KEY, "objective") + + +def register_user_attr_form_widgets( + study: optuna.Study, widgets: list[ObjectiveFormWidget] +) -> None: + widget_dicts = [w.to_dict() for w in widgets] + study._storage.set_study_system_attr(study._study_id, SYSTEM_ATTR_KEY, widget_dicts) + study._storage.set_study_system_attr(study._study_id, SYSTEM_ATTR_OUTPUT_TYPE_KEY, "user_attr") def get_objective_form_widgets_json( diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index 02170570..b5ef1024 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -2,6 +2,7 @@ from __future__ import annotations import json from typing import Any +from typing import Optional from typing import TYPE_CHECKING from typing import Union @@ -121,6 +122,7 @@ def serialize_study_detail( union: list[tuple[str, BaseDistribution]], union_user_attrs: list[tuple[str, bool]], has_intermediate_values: bool, + form_widgets_output_type: Optional[str], ) -> dict[str, Any]: serialized: dict[str, Any] = { "name": summary.study_name, @@ -147,6 +149,7 @@ def serialize_study_detail( objective_form_widgets = get_objective_form_widgets_json(system_attrs) if objective_form_widgets: serialized["objective_form_widgets"] = objective_form_widgets + serialized["form_widgets_output_type"] = form_widgets_output_type return serialized diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index 2cd32c53..849bf287 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -68,6 +68,7 @@ interface StudyDetailResponse { note: Note objective_names?: string[] objective_form_widgets?: ObjectiveFormWidget[] + form_widgets_output_type?: string } export const getStudyDetailAPI = ( @@ -101,6 +102,7 @@ export const getStudyDetailAPI = ( note: res.data.note, objective_names: res.data.objective_names, objective_form_widgets: res.data.objective_form_widgets, + form_widgets_output_type: res.data.form_widgets_output_type, } }) } diff --git a/optuna_dashboard/ts/components/ObjectiveForm.tsx b/optuna_dashboard/ts/components/ObjectiveForm.tsx index 6433bbe7..704a7916 100644 --- a/optuna_dashboard/ts/components/ObjectiveForm.tsx +++ b/optuna_dashboard/ts/components/ObjectiveForm.tsx @@ -21,7 +21,8 @@ export const ObjectiveForm: FC<{ directions: StudyDirection[] names: string[] widgets: ObjectiveFormWidget[] -}> = ({ trial, directions, names, widgets }) => { + outputType: string +}> = ({ trial, directions, names, widgets, outputType }) => { const theme = useTheme() const action = actionCreator() const [values, setValues] = useState<(number | null)[]>( @@ -64,8 +65,16 @@ export const ObjectiveForm: FC<{ const handleSubmit = (e: React.MouseEvent): void => { e.preventDefault() - const user_attrs = Object.fromEntries(widgets.map((widget, i) => [widget.description, values[i]])) - action.saveTrialUserAttrs(trial.study_id, trial.trial_id, user_attrs) + if (outputType == "objective") { + const filtered = values.filter((v): v is number => v !== null) + if (filtered.length !== directions.length) { + return + } + action.tellTrial(trial.study_id, trial.trial_id, "Complete", filtered) + } else if (outputType == "user_attr") { + const user_attrs = Object.fromEntries(widgets.map((widget, i) => [widget.description, values[i]])) + action.saveTrialUserAttrs(trial.study_id, trial.trial_id, user_attrs) + } } const getObjectiveName = (i: number): string => { diff --git a/optuna_dashboard/ts/components/TrialList.tsx b/optuna_dashboard/ts/components/TrialList.tsx index 7ef89d73..85ba3d68 100644 --- a/optuna_dashboard/ts/components/TrialList.tsx +++ b/optuna_dashboard/ts/components/TrialList.tsx @@ -144,12 +144,14 @@ const TrialListDetail: FC<{ directions: StudyDirection[] objectiveNames: string[] objectiveFormWidgets: ObjectiveFormWidget[] + formWigetsOutputType: string }> = ({ trial, isBestTrial, directions, objectiveNames, objectiveFormWidgets, + formWigetsOutputType }) => { const theme = useTheme() const artifactEnabled = useRecoilValue(artifactIsAvailable) @@ -297,6 +299,7 @@ const TrialListDetail: FC<{ directions={directions} names={objectiveNames} widgets={objectiveFormWidgets} + outputType={formWigetsOutputType} /> )} {trial.state === "Complete" && directions.length > 0 && ( @@ -816,6 +819,9 @@ export const TrialList: FC<{ studyDetail: StudyDetail | null }> = ({ objectiveFormWidgets={ studyDetail?.objective_form_widgets || [] } + formWigetsOutputType={ + studyDetail?.form_widgets_output_type || "" + } /> ))} diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index 1d30bc57..87b33a29 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -176,6 +176,7 @@ type StudyDetail = { note: Note objective_names?: string[] objective_form_widgets?: ObjectiveFormWidget[] + form_widgets_output_type?: string } type StudyDetails = {