From ec5c46bc6cf01a88fd964b1044251bdfe53f03cc Mon Sep 17 00:00:00 2001 From: Gene Date: Mon, 4 Dec 2023 19:18:17 +0800 Subject: [PATCH] [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)}")