mirror of
https://github.com/wassname/optuna-dashboard.git
synced 2026-10-04 12:50:44 +08:00
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:
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)
|
||||
Reference in new issue
Block a user