diff --git a/optuna_dashboard/_app.py b/optuna_dashboard/_app.py index 1c9e987a..98571add 100644 --- a/optuna_dashboard/_app.py +++ b/optuna_dashboard/_app.py @@ -297,14 +297,20 @@ def create_app(storage: BaseStorage, debug: bool = False) -> Bottle: response.status = 404 # Not found return {"reason": f"study_id={study_id} is not found"} trials = get_trials(storage, study_id) - intersection, union, has_intermeridate_values = get_cached_extra_study_property( - study_id, trials - ) + ( + # TODO: intersection_search_space and union_search_space look more clear since now we + # have union_user_attrs. + intersection, + union, + union_user_attrs, + has_intermeridate_values, + ) = get_cached_extra_study_property(study_id, trials) return serialize_study_detail( summary, trials[after:], intersection, union, + union_user_attrs, has_intermeridate_values, ) diff --git a/optuna_dashboard/_cached_extra_study_property.py b/optuna_dashboard/_cached_extra_study_property.py index 3688a010..316e6d61 100644 --- a/optuna_dashboard/_cached_extra_study_property.py +++ b/optuna_dashboard/_cached_extra_study_property.py @@ -23,7 +23,7 @@ states_of_interest = [TrialState.COMPLETE, TrialState.PRUNED] def get_cached_extra_study_property( study_id: int, trials: List[FrozenTrial] -) -> Tuple[SearchSpaceListT, SearchSpaceListT, bool]: +) -> Tuple[SearchSpaceListT, SearchSpaceListT, List[str], bool]: with cached_extra_study_property_cache_lock: cached_extra_study_property = cached_extra_study_property_cache.get(study_id, None) if cached_extra_study_property is None: @@ -33,6 +33,7 @@ def get_cached_extra_study_property( return ( cached_extra_study_property.intersection, cached_extra_study_property.union, + cached_extra_study_property.union_user_attrs, cached_extra_study_property.has_intermediate_values, ) @@ -40,8 +41,11 @@ def get_cached_extra_study_property( class _CachedExtraStudyProperty: def __init__(self) -> None: self._cursor: int = -1 + # TODO: intersection_search_space and union_search_space look more clear since now we have + # union_user_attrs. self._intersection: Optional[SearchSpaceSetT] = None self._union: SearchSpaceSetT = set() + self._union_user_attrs: Set[str] = set() self.has_intermediate_values: bool = False @property @@ -58,6 +62,12 @@ class _CachedExtraStudyProperty: union.sort(key=lambda x: x[0]) return union + @property + def union_user_attrs(self) -> List[str]: + union = list(self._union_user_attrs) + union.sort() + return union + def update(self, trials: List[FrozenTrial]) -> None: next_cursor = self._cursor for trial in reversed(trials): @@ -81,4 +91,7 @@ class _CachedExtraStudyProperty: else: self._intersection = self._intersection.intersection(current) + current_user_attrs = set(n for n in trial.user_attrs.keys()) + self._union_user_attrs = self._union_user_attrs.union(current_user_attrs) + self._cursor = next_cursor diff --git a/optuna_dashboard/_serializer.py b/optuna_dashboard/_serializer.py index dc754e82..9dca618e 100644 --- a/optuna_dashboard/_serializer.py +++ b/optuna_dashboard/_serializer.py @@ -71,6 +71,7 @@ def serialize_study_detail( trials: List[FrozenTrial], intersection: List[Tuple[str, BaseDistribution]], union: List[Tuple[str, BaseDistribution]], + union_user_attrs: List[str], has_intermediate_values: bool, ) -> Dict[str, Any]: serialized: Dict[str, Any] = { @@ -83,6 +84,7 @@ def serialize_study_detail( serialized["trials"] = [serialize_frozen_trial(summary._study_id, trial) for trial in trials] serialized["intersection_search_space"] = serialize_search_space(intersection) serialized["union_search_space"] = serialize_search_space(union) + serialized["union_user_attrs"] = union_user_attrs serialized["has_intermediate_values"] = has_intermediate_values serialized["note"] = note.get_note_from_system_attrs(summary.system_attrs) return serialized diff --git a/optuna_dashboard/ts/apiClient.ts b/optuna_dashboard/ts/apiClient.ts index d0f30038..2f00395b 100644 --- a/optuna_dashboard/ts/apiClient.ts +++ b/optuna_dashboard/ts/apiClient.ts @@ -44,6 +44,7 @@ interface StudyDetailResponse { trials: TrialResponse[] intersection_search_space: SearchSpace[] union_search_space: SearchSpace[] + union_user_attrs: string[] has_intermediate_values: boolean note: { version: number @@ -73,6 +74,7 @@ export const getStudyDetailAPI = ( trials: trials, union_search_space: res.data.union_search_space, intersection_search_space: res.data.intersection_search_space, + union_user_attrs: res.data.union_user_attrs, has_intermediate_values: res.data.has_intermediate_values, note: res.data.note, } diff --git a/optuna_dashboard/ts/components/StudyDetail.tsx b/optuna_dashboard/ts/components/StudyDetail.tsx index feb9a88c..2a52e47e 100644 --- a/optuna_dashboard/ts/components/StudyDetail.tsx +++ b/optuna_dashboard/ts/components/StudyDetail.tsx @@ -650,6 +650,36 @@ export const TrialTable: FC<{ studyDetail: StudyDetail | null }> = ({ }) } + studyDetail?.union_user_attrs.forEach((attr_name) => { + columns.push({ + field: "user_attrs", + label: `User attribute ${attr_name}`, + toCellValue: (i) => + trials[i].user_attrs.find((attr) => attr.key === attr_name)?.value || + null, + sortable: true, + filterable: true, + less: (firstEl, secondEl): number => { + const firstVal = firstEl.user_attrs.find( + (attr) => attr.key === attr_name + )?.value + const secondVal = secondEl.user_attrs.find( + (attr) => attr.key === attr_name + )?.value + + if (firstVal === secondVal) { + return 0 + } else if (firstVal && secondVal) { + return firstVal < secondVal ? 1 : -1 + } else if (firstVal) { + return -1 + } else { + return 1 + } + }, + }) + }) + const collapseParamColumns: DataGridColumn[] = [ { field: "name", label: "Name", sortable: true }, { field: "value", label: "Value", sortable: true }, diff --git a/optuna_dashboard/ts/types/index.d.ts b/optuna_dashboard/ts/types/index.d.ts index 951e858e..2f15ea75 100644 --- a/optuna_dashboard/ts/types/index.d.ts +++ b/optuna_dashboard/ts/types/index.d.ts @@ -84,6 +84,7 @@ declare interface StudyDetail { trials: Trial[] intersection_search_space: SearchSpace[] union_search_space: SearchSpace[] + union_user_attrs: string[] has_intermediate_values: boolean note: Note } diff --git a/typescript_tests/TrialTable.test.tsx b/typescript_tests/TrialTable.test.tsx index da783fec..be9dd663 100644 --- a/typescript_tests/TrialTable.test.tsx +++ b/typescript_tests/TrialTable.test.tsx @@ -74,6 +74,7 @@ const studyDetail = { attributes: { low: -3, high: 3 }, }, ], + union_user_attrs: ["foo", "bar"], has_intermediate_values: false, note: { version: 0,