diff --git a/learnware/client/container.py b/learnware/client/container.py index 15236b0..8e7573a 100644 --- a/learnware/client/container.py +++ b/learnware/client/container.py @@ -31,7 +31,7 @@ class ModelEnvContainer(BaseModel): with open(model_path, "wb") as model_fp: pickle.dump(self.model_config, model_fp) - + system_execute( [ "conda", diff --git a/learnware/client/learnware_client.py b/learnware/client/learnware_client.py index 3ec9bcf..1c6cfed 100644 --- a/learnware/client/learnware_client.py +++ b/learnware/client/learnware_client.py @@ -337,11 +337,8 @@ class LearnwareClient: if load_model: learnware_obj.instantiate_model() - pass - + return learnware_obj - pass - pass def system(self, command): retcd = os.system(command) diff --git a/tests/test_client/test_download.py b/tests/test_client/test_download.py index 5dba8ae..7314ac8 100644 --- a/tests/test_client/test_download.py +++ b/tests/test_client/test_download.py @@ -1,12 +1,55 @@ import os +import zipfile import numpy as np import learnware +from learnware.learnware import get_learnware_from_dirpath from learnware.client import LearnwareClient from learnware.client.container import ModelEnvContainer, LearnwaresContainer from learnware.learnware.reuse import AveragingReuser +def test_container(zip_paths): + semantic_specification = dict() + semantic_specification["Data"] = {"Type": "Class", "Values": ["Text"]} + semantic_specification["Task"] = {"Type": "Class", "Values": ["Ranking"]} + semantic_specification["Library"] = {"Type": "Class", "Values": ["Scikit-learn"]} + semantic_specification["Scenario"] = {"Type": "Tag", "Values": "Financial"} + semantic_specification["Name"] = {"Type": "String", "Values": "test"} + semantic_specification["Description"] = {"Type": "String", "Values": "test"} + + learnware_list = [] + for id, zip_path in enumerate(zip_paths): + dir_path = zip_path[:-4] + with zipfile.ZipFile(zip_path, "r") as z_file: + z_file.extractall(dir_path) + + learnware = get_learnware_from_dirpath(f"test_id{id}", semantic_specification, dir_path) + learnware_list.append(learnware) + + with LearnwaresContainer(learnware_list, zip_paths) as env_container: + learnware_list = env_container.get_learnware_list_with_container() + reuser = AveragingReuser(learnware_list, mode="vote_by_label") + input_array = np.random.random(size=(20, 13)) + print(reuser.predict(input_array)) + + for idx, learnware in enumerate(learnware_list): + print(f"learnware_{idx}", learnware.predict(input_array)) + + +def test_load(zip_paths): + learnware_list = [client.load_learnware(file, load_model=False) for file in zip_paths] + + with LearnwaresContainer(learnware_list, zip_paths) as env_container: + learnware_list = env_container.get_learnware_list_with_container() + reuser = AveragingReuser(learnware_list, mode="vote_by_label") + input_array = np.random.random(size=(20, 13)) + print(reuser.predict(input_array)) + + for idx, learnware in enumerate(learnware_list): + print(f"learnware_{idx}", learnware.predict(input_array)) + + if __name__ == "__main__": email = "liujd@lamda.nju.edu.cn" token = "f7e647146a314c6e8b4e2e1079c4bca4" @@ -21,13 +64,5 @@ if __name__ == "__main__": zip_paths[i] = os.path.join(root, zip_paths[i]) client.download_learnware(learnware_ids[i], zip_paths[i]) - learnware_list = [client.load_learnware(file, load_model=False) for file in zip_paths] - - with LearnwaresContainer(learnware_list, zip_paths) as env_container: - learnware_list = env_container.get_learnware_list_with_container() - reuser = AveragingReuser(learnware_list, mode="vote_by_label") - input_array = np.random.random(size=(20, 13)) - print(reuser.predict(input_array)) - - for idx, learnware in enumerate(learnware_list): - print(f"learnware_{idx}", reuser.predict(learnware)) \ No newline at end of file + test_container(zip_paths) + # test_load(zip_paths) \ No newline at end of file