From 4228012194b2b786196ca8ab0d4eb28742fc1b04 Mon Sep 17 00:00:00 2001 From: Gene Date: Mon, 4 Dec 2023 19:17:44 +0800 Subject: [PATCH 1/3] [MNT] add dist.isfinite check --- learnware/market/easy/checker.py | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/learnware/market/easy/checker.py b/learnware/market/easy/checker.py index 57cae22..a7cea3b 100644 --- a/learnware/market/easy/checker.py +++ b/learnware/market/easy/checker.py @@ -115,7 +115,11 @@ class EasyStatChecker(BaseChecker): # Check if statistical specification is computable in dist() stat_spec = learnware.get_specification().get_stat_spec_by_name(spec_type) - stat_spec.dist(stat_spec) + distance = stat_spec.dist(stat_spec) + if not np.isfinite(distance): + message = f"The distance between statistical specifications is not finite, where distance={distance}" + logger.warning(message) + return self.INVALID_LEARNWARE, message if spec_type == "RKMETableSpecification": if not isinstance(input_shape, tuple) or not all(isinstance(item, int) for item in input_shape): From ec5c46bc6cf01a88fd964b1044251bdfe53f03cc Mon Sep 17 00:00:00 2001 From: Gene Date: Mon, 4 Dec 2023 19:18:17 +0800 Subject: [PATCH 2/3] [MNT] filter learnware with infinite dist --- learnware/market/easy/searcher.py | 21 +++++++++++++++------ 1 file changed, 15 insertions(+), 6 deletions(-) diff --git a/learnware/market/easy/searcher.py b/learnware/market/easy/searcher.py index 48bbfa8..f08b9b2 100644 --- a/learnware/market/easy/searcher.py +++ b/learnware/market/easy/searcher.py @@ -538,14 +538,20 @@ class EasyStatSearcher(BaseSearcher): both lists are sorted by mmd dist """ rkme_list = [learnware.specification.get_stat_spec_by_name(self.stat_spec_type) for learnware in learnware_list] - mmd_dist_list = [] - for rkme in rkme_list: - mmd_dist = rkme.dist(user_rkme) - mmd_dist_list.append(mmd_dist) + filtered_idx_list, mmd_dist_list = [], [] + for idx in range(len(rkme_list)): + mmd_dist = rkme_list[idx].dist(user_rkme) + if np.isfinite(mmd_dist): + mmd_dist_list.append(mmd_dist) + filtered_idx_list.append(idx) + else: + logger.warning( + f"The distance between user_spec and learnware_spec (id: {learnware_list[idx].id}) is not finite, where distance is {mmd_dist}" + ) - sorted_idx_list = sorted(range(len(learnware_list)), key=lambda k: mmd_dist_list[k]) + sorted_idx_list = sorted(range(len(mmd_dist_list)), key=lambda k: mmd_dist_list[k]) sorted_dist_list = [mmd_dist_list[idx] for idx in sorted_idx_list] - sorted_learnware_list = [learnware_list[idx] for idx in sorted_idx_list] + sorted_learnware_list = [learnware_list[filtered_idx_list[idx]] for idx in sorted_idx_list] return sorted_dist_list, sorted_learnware_list @@ -561,6 +567,9 @@ class EasyStatSearcher(BaseSearcher): raise KeyError("No supported stat specification is given in the user info") user_rkme = user_info.stat_info[self.stat_spec_type] + if not np.isfinite(user_rkme.dist(user_rkme)): + raise ValueError("The distance between uploaded statistical specifications is not finite!") + learnware_list = self._filter_by_rkme_spec_metadata(learnware_list, user_rkme) logger.info(f"After filter by rkme dimension, learnware_list length is {len(learnware_list)}") From 9bba4146f1033230f8de724d17cbafebe27c6c81 Mon Sep 17 00:00:00 2001 From: Gene Date: Mon, 4 Dec 2023 19:23:22 +0800 Subject: [PATCH 3/3] [MNT] modify details --- learnware/market/easy/checker.py | 2 +- learnware/market/easy/searcher.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) diff --git a/learnware/market/easy/checker.py b/learnware/market/easy/checker.py index a7cea3b..95c0f1a 100644 --- a/learnware/market/easy/checker.py +++ b/learnware/market/easy/checker.py @@ -115,7 +115,7 @@ class EasyStatChecker(BaseChecker): # Check if statistical specification is computable in dist() stat_spec = learnware.get_specification().get_stat_spec_by_name(spec_type) - distance = stat_spec.dist(stat_spec) + distance = float(stat_spec.dist(stat_spec)) if not np.isfinite(distance): message = f"The distance between statistical specifications is not finite, where distance={distance}" logger.warning(message) diff --git a/learnware/market/easy/searcher.py b/learnware/market/easy/searcher.py index f08b9b2..7f6690f 100644 --- a/learnware/market/easy/searcher.py +++ b/learnware/market/easy/searcher.py @@ -540,7 +540,7 @@ class EasyStatSearcher(BaseSearcher): rkme_list = [learnware.specification.get_stat_spec_by_name(self.stat_spec_type) for learnware in learnware_list] filtered_idx_list, mmd_dist_list = [], [] for idx in range(len(rkme_list)): - mmd_dist = rkme_list[idx].dist(user_rkme) + mmd_dist = float(rkme_list[idx].dist(user_rkme)) if np.isfinite(mmd_dist): mmd_dist_list.append(mmd_dist) filtered_idx_list.append(idx) @@ -567,7 +567,7 @@ class EasyStatSearcher(BaseSearcher): raise KeyError("No supported stat specification is given in the user info") user_rkme = user_info.stat_info[self.stat_spec_type] - if not np.isfinite(user_rkme.dist(user_rkme)): + if not np.isfinite(float(user_rkme.dist(user_rkme))): raise ValueError("The distance between uploaded statistical specifications is not finite!") learnware_list = self._filter_by_rkme_spec_metadata(learnware_list, user_rkme)