Merge pull request #302 from c-bata/add-tests-for-cached-estra-study-property

Fix `union_user_attrs` property of `_CachedExtraStudyProperty`.
This commit is contained in:
Masashi Shibata authored and GitHub committed 2022-12-08 12:50:25 +09:00
commit 4a0ba0961d
2 files changed
+70 -6

No files matched your search

@@ -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
@@ -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)