diff --git a/examples/example_m5/main.py b/examples/example_m5/main.py index 82ff2d3..1d3a39b 100644 --- a/examples/example_m5/main.py +++ b/examples/example_m5/main.py @@ -101,7 +101,7 @@ class M5DatasetWorkflow: m5 = DataLoader() idx_list = m5.get_idx_list() - algo_list = ['lgb'] # algo_list = ["ridge", "lgb"] + algo_list = ["lgb"] # algo_list = ["ridge", "lgb"] curr_root = os.path.dirname(os.path.abspath(__file__)) curr_root = os.path.join(curr_root, "learnware_pool") diff --git a/examples/example_m5/upload.py b/examples/example_m5/upload.py index ba0234a..0c9e209 100644 --- a/examples/example_m5/upload.py +++ b/examples/example_m5/upload.py @@ -65,7 +65,7 @@ def main(): "Device": {"Values": ["CPU"], "Type": "Tag"}, "Scenario": {"Values": ["Business"], "Type": "Tag"}, "Description": {"Values": "A sales-forecasting model from Walmart store", "Type": "String"}, - "Name": {"Values": {name}, "Type": "String"}, + "Name": {"Values": name, "Type": "String"}, } res = session.post( submit_url, diff --git a/examples/example_pfs/upload.py b/examples/example_pfs/upload.py index 04dbcd8..ed8449f 100644 --- a/examples/example_pfs/upload.py +++ b/examples/example_pfs/upload.py @@ -68,7 +68,7 @@ def main(): "Values": "A sales-forecasting model from Predict Future Sales Competition on Kaggle", "Type": "String", }, - "Name": {"Values": {name}, "Type": "String"}, + "Name": {"Values": name, "Type": "String"}, } res = session.post( submit_url, diff --git a/learnware/market/easy.py b/learnware/market/easy.py index 130d918..0133fa1 100644 --- a/learnware/market/easy.py +++ b/learnware/market/easy.py @@ -333,7 +333,7 @@ class EasyMarket(BaseMarket): learnware_list: List[Learnware], user_rkme: RKMEStatSpecification, max_search_num: int, - weight_cutoff: float = 0.95 + weight_cutoff: float = 0.95, ) -> Tuple[List[float], List[Learnware]]: """Select learnwares based on a total mixture ratio, then recalculate their mixture weights @@ -372,15 +372,15 @@ class EasyMarket(BaseMarket): mixture_list.append(learnware_list[idx]) else: break - + if len(mixture_list) <= 1: mixture_list = [learnware_list[sort_by_weight_idx_list[0]]] mixture_weight = [1] else: if len(mixture_list) > max_search_num: - mixture_list = mixture_list[:max_search_num] + mixture_list = mixture_list[:max_search_num] mixture_weight, _ = self._calculate_rkme_spec_mixture_weight(mixture_list, user_rkme) - + return mixture_weight, mixture_list def _filter_by_rkme_spec_single( @@ -618,11 +618,11 @@ class EasyMarket(BaseMarket): sorted_score_list, single_learnware_list = self._filter_by_rkme_spec_single( sorted_score_list, single_learnware_list ) - if search_method == 'auto': + if search_method == "auto": weight_list, mixture_learnware_list = self._search_by_rkme_spec_mixture_auto( learnware_list, user_rkme, max_search_num ) - elif search_method == 'greedy': + elif search_method == "greedy": weight_list, mixture_learnware_list = self._search_by_rkme_spec_mixture_greedy( learnware_list, user_rkme, max_search_num )