diff --git a/optuna_dashboard/_cached_extra_study_property.py b/optuna_dashboard/_cached_extra_study_property.py index 316e6d61..1e9de3fc 100644 --- a/optuna_dashboard/_cached_extra_study_property.py +++ b/optuna_dashboard/_cached_extra_study_property.py @@ -18,8 +18,6 @@ SearchSpaceListT = List[Tuple[str, BaseDistribution]] cached_extra_study_property_cache_lock = threading.Lock() cached_extra_study_property_cache: Dict[int, "_CachedExtraStudyProperty"] = {} -states_of_interest = [TrialState.COMPLETE, TrialState.PRUNED] - def get_cached_extra_study_property( study_id: int, trials: List[FrozenTrial] @@ -77,7 +75,9 @@ class _CachedExtraStudyProperty: if not trial.state.is_finished(): next_cursor = trial.number - if trial.state not in states_of_interest: + current_user_attrs = set(n for n in trial.user_attrs.keys()) + self._union_user_attrs = self._union_user_attrs.union(current_user_attrs) + if trial.state == TrialState.FAIL: continue if not self.has_intermediate_values and len(trial.intermediate_values) > 0: @@ -91,7 +91,4 @@ 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/python_tests/test_cached_extra_study_property.py b/python_tests/test_cached_extra_study_property.py index 08280f59..8ded533b 100644 --- a/python_tests/test_cached_extra_study_property.py +++ b/python_tests/test_cached_extra_study_property.py @@ -117,6 +117,30 @@ class _CachedExtraStudyPropertySearchSpaceTestCase(TestCase): self.assertEqual(len(cached_extra_study_property.intersection), 0) self.assertEqual(len(cached_extra_study_property.union), 3) + def test_contains_failed_trials(self) -> None: + distributions = { + "x0": UniformDistribution(low=0, high=10), + "x1": UniformDistribution(low=0, high=10), + } + params = { + "x0": 0.5, + "x1": 0.5, + } + trials = [ + create_trial( + state=TrialState.COMPLETE, value=0, distributions=distributions, params=params + ), + create_trial(state=TrialState.FAIL, value=0, distributions={}, params={}), + create_trial( + state=TrialState.COMPLETE, value=0, distributions=distributions, params=params + ), + ] + cached_extra_study_property = _CachedExtraStudyProperty() + cached_extra_study_property.update(trials) + + self.assertEqual(len(cached_extra_study_property.intersection), 2) + self.assertEqual(len(cached_extra_study_property.union), 2) + class _CachedExtraStudyPropertyIntermediateTestCase(TestCase): def setUp(self) -> None: @@ -183,3 +207,46 @@ class _CachedExtraStudyPropertyIntermediateTestCase(TestCase): cached_extra_study_property = _CachedExtraStudyProperty() cached_extra_study_property.update(trials) self.assertFalse(cached_extra_study_property.has_intermediate_values) + + +class _CachedExtraStudyPropertyUserAttrs(TestCase): + def setUp(self) -> None: + optuna.logging.set_verbosity(optuna.logging.ERROR) + warnings.simplefilter("ignore", category=ExperimentalWarning) + + def test_contains_failed_trials(self) -> None: + distributions = { + "x0": UniformDistribution(low=0, high=10), + "x1": UniformDistribution(low=0, high=10), + } + params = { + "x0": 0.5, + "x1": 0.5, + } + trials = [ + create_trial( + state=TrialState.COMPLETE, + value=0, + distributions=distributions, + params=params, + user_attrs={"foo": "foo"}, + ), + create_trial( + state=TrialState.FAIL, + value=0, + distributions={}, + params={}, + user_attrs={"bar": "bar"}, + ), + create_trial( + state=TrialState.COMPLETE, + value=0, + distributions=distributions, + params=params, + user_attrs={"baz": "baz"}, + ), + ] + cached_extra_study_property = _CachedExtraStudyProperty() + cached_extra_study_property.update(trials) + + self.assertEqual(len(cached_extra_study_property.union_user_attrs), 3)