Browse Source

[MNT] Modify upload in example

tags/v0.3.2
Gene 3 years ago
parent
commit
20b9613173
4 changed files with 9 additions and 9 deletions
  1. +1
    -1
      examples/example_m5/main.py
  2. +1
    -1
      examples/example_m5/upload.py
  3. +1
    -1
      examples/example_pfs/upload.py
  4. +6
    -6
      learnware/market/easy.py

+ 1
- 1
examples/example_m5/main.py View File

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


+ 1
- 1
examples/example_m5/upload.py View File

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


+ 1
- 1
examples/example_pfs/upload.py View File

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


+ 6
- 6
learnware/market/easy.py View File

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


Loading…
Cancel
Save