diff --git a/examples/example_pfs/main.py b/examples/example_pfs/main.py index 68301b1..ae20c98 100644 --- a/examples/example_pfs/main.py +++ b/examples/example_pfs/main.py @@ -1,13 +1,14 @@ import os import fire import zipfile +import numpy as np from tqdm import tqdm from shutil import copyfile, rmtree import learnware from learnware.market import EasyMarket, BaseUserInfo from learnware.market import database_ops -from learnware.learnware import Learnware, JobSelectorReuser +from learnware.learnware import Learnware, JobSelectorReuser, EnsembleReuser import learnware.specification as specification from pfs import Dataloader @@ -112,8 +113,8 @@ class PFSDatasetWorkflow: rmtree(dir_path) def test(self, regenerate_flag=False): - # self.prepare_learnware(regenerate_flag) - # self._init_learnware_market() + self.prepare_learnware(regenerate_flag) + self._init_learnware_market() easy_market = EasyMarket() print("Total Item:", len(easy_market)) @@ -150,17 +151,32 @@ class PFSDatasetWorkflow: for score, learnware in zip(sorted_score_list, single_learnware_list): pred_y = learnware.predict(test_x) loss_list.append(pfs.score(test_y, pred_y)) - print(f"Top1-score: {sorted_score_list[0]}, learnware_id: {learnware.id}, loss: {loss_list[-1]}") + print( + f"Top1-score: {sorted_score_list[0]}, learnware_id: {single_learnware_list[0].id}, loss: {loss_list[-1]}" + ) mixture_id = " ".join([learnware.id for learnware in mixture_learnware_list]) print(f"mixture_score: {mixture_score}, mixture_learnware: {mixture_id}") - reuse_baseline = JobSelectorReuser(learnware_list=mixture_learnware_list) - reuse_predict = reuse_baseline.predict(user_data=test_x) - reuse_score = pfs.score(test_y, reuse_predict) - print(f"mixture reuse loss: {reuse_score}\n") - - sinle_score_list.append() + reuse_job_selector = JobSelectorReuser(learnware_list=mixture_learnware_list) + job_selector_predict_y = reuse_job_selector.predict(user_data=test_x) + job_selector_score = pfs.score(test_y, job_selector_predict_y) + print(f"mixture reuse loss (job selector): {job_selector_score}") + + reuse_ensemble = EnsembleReuser(learnware_list=mixture_learnware_list) + ensemble_predict_y = reuse_ensemble.predict(user_data=test_x) + ensemble_score = pfs.score(test_y, ensemble_predict_y) + print(f"mixture reuse loss (ensemble): {ensemble_score}\n") + + sinle_score_list.append(loss_list[0]) + random_score_list.append(np.mean(loss_list)) + job_selector_score_list.append(job_selector_score) + ensemble_score_list.append(ensemble_score) + + print(f"Single search score: {np.mean(sinle_score_list)}") + print(f"Job selector score: {np.mean(job_selector_score_list)}") + print(f"Average ensemble score: {np.mean(ensemble_score_list)}") + print(f"Random search score: {np.mean(random_score_list)}") if __name__ == "__main__": diff --git a/learnware/learnware/__init__.py b/learnware/learnware/__init__.py index b72a110..d1cf547 100644 --- a/learnware/learnware/__init__.py +++ b/learnware/learnware/__init__.py @@ -2,7 +2,7 @@ import os import copy from .base import Learnware, BaseReuser -from .reuse import JobSelectorReuser +from .reuse import JobSelectorReuser, EnsembleReuser from .utils import get_stat_spec_from_config, get_model_from_config from ..specification import Specification diff --git a/learnware/learnware/reuse.py b/learnware/learnware/reuse.py index 58c3914..0b49f10 100644 --- a/learnware/learnware/reuse.py +++ b/learnware/learnware/reuse.py @@ -223,3 +223,44 @@ class JobSelectorReuser(BaseReuser): ) return model + + +class EnsembleReuser(BaseReuser): + """Baseline Multiple Learnware Reuser uing Ensemble Method""" + + def __init__(self, learnware_list: List[Learnware]): + """The initialization method for ensemble reuser + + Parameters + ---------- + learnware_list : List[Learnware] + The learnware list, which should have RKME Specification for each learnweare + """ + super(EnsembleReuser, self).__init__(learnware_list) + + def predict(self, user_data: np.ndarray) -> np.ndarray: + """Give prediction for user data using baseline ensemble method + + Parameters + ---------- + user_data : np.ndarray + User's labeled raw data. + + Returns + ------- + np.ndarray + Prediction given by ensemble method + """ + mean_pred_y = None + + for idx in range(len(self.learnware_list)): + pred_y = self.learnware_list[idx].predict(user_data) + + if mean_pred_y is None: + mean_pred_y = pred_y + else: + mean_pred_y += pred_y + + mean_pred_y /= len(self.learnware_list) + + return mean_pred_y