From e062a181cff4c0ce96b97a9c0e18b6c0651ca169 Mon Sep 17 00:00:00 2001 From: Gene Date: Fri, 13 Oct 2023 15:59:12 +0800 Subject: [PATCH] [MNT] format code --- learnware/client/container.py | 8 ++++---- learnware/client/learnware_client.py | 21 +++++++++++++-------- tests/test_client/test_load.py | 22 +++++++++++++--------- 3 files changed, 30 insertions(+), 21 deletions(-) diff --git a/learnware/client/container.py b/learnware/client/container.py index c924722..2221523 100644 --- a/learnware/client/container.py +++ b/learnware/client/container.py @@ -127,11 +127,11 @@ class LearnwaresContainer: ) for _learnware, _zippath in zip(learnware_list, learnware_zippaths) ] - + model_list = [_learnware.get_model() for _learnware in self.learnware_list] with ProcessPoolExecutor(max_workers=max(os.cpu_count() // 2, 1)) as executor: executor.map(self._initialize_model_container, model_list) - + atexit.register(self.cleanup) @staticmethod @@ -144,7 +144,7 @@ class LearnwaresContainer: def get_learnware_list_with_container(self): return self.learnware_list - + def cleanup(self): for _learnware in self.learnware_list: - self._destroy_model_container(_learnware.get_model()) \ No newline at end of file + self._destroy_model_container(_learnware.get_model()) diff --git a/learnware/client/learnware_client.py b/learnware/client/learnware_client.py index d3d79c9..8f07408 100644 --- a/learnware/client/learnware_client.py +++ b/learnware/client/learnware_client.py @@ -310,7 +310,12 @@ class LearnwareClient: return semantic_conf[key.value]["Values"] - def load_learnware(self, learnware_path: Union[str, List[str]] = None, learnware_id: Union[str, List[str]] = None, runnable_option: str = None): + def load_learnware( + self, + learnware_path: Union[str, List[str]] = None, + learnware_id: Union[str, List[str]] = None, + runnable_option: str = None, + ): """Load learnware by learnware zip file or learnware id (zip file has higher priority) Parameters @@ -334,14 +339,14 @@ class LearnwareClient: if learnware_path is None and learnware_id is None: raise ValueError("Requires one of learnware_path or learnware_id") - + def _get_learnware_by_id(_learnware_id): self.tempdir_list.append(tempfile.TemporaryDirectory(prefix="learnware_")) tempdir = self.tempdir_list[-1].name zip_path = os.path.join(tempdir, f"{str(uuid.uuid4())}.zip") self.download_learnware(_learnware_id, zip_path) return zip_path, _get_learnware_by_path(zip_path, tempdir=tempdir) - + def _get_learnware_by_path(_learnware_zippath, tempdir=None): if tempdir is None: self.tempdir_list.append(tempfile.TemporaryDirectory(prefix="learnware_")) @@ -368,7 +373,7 @@ class LearnwareClient: semantic_specification = json.load(fin) return learnware.get_learnware_from_dirpath(learnware_id, semantic_specification, tempdir) - + learnware_list = [] zip_paths = [] if learnware_path is not None: @@ -376,7 +381,7 @@ class LearnwareClient: zip_paths = [learnware_path] elif isinstance(learnware_path, list): zip_paths = learnware_path - + for zip_path in zip_paths: learnware_obj = _get_learnware_by_path(zip_path) learnware_list.append(learnware_obj) @@ -385,12 +390,12 @@ class LearnwareClient: id_list = [learnware_id] elif isinstance(learnware_id, list): id_list = learnware_id - + for idx in id_list: zip_path, learnware_obj = _get_learnware_by_id(idx) zip_paths.append(zip_path) learnware_list.append(learnware_obj) - + if runnable_option is not None: if runnable_option == "normal": for i in range(len(learnware_list)): @@ -398,7 +403,7 @@ class LearnwareClient: elif runnable_option == "conda_env": env_container = LearnwaresContainer(learnware_list, zip_paths) learnware_list = env_container.get_learnware_list_with_container() - + if len(learnware_list) == 1: return learnware_list[0] else: diff --git a/tests/test_client/test_load.py b/tests/test_client/test_load.py index 1706e00..67981dc 100644 --- a/tests/test_client/test_load.py +++ b/tests/test_client/test_load.py @@ -11,7 +11,6 @@ from learnware.learnware.reuse import AveragingReuser class TestLearnwareLoad(unittest.TestCase): - def setUp(self): unittest.TestCase.setUpClass() email = "liujd@lamda.nju.edu.cn" @@ -27,8 +26,11 @@ class TestLearnwareLoad(unittest.TestCase): def test_load_single_learnware_by_zippath(self): for (learnware_id, zip_path) in zip(self.learnware_ids, self.zip_paths): self.client.download_learnware(learnware_id, zip_path) - - learnware_list = [self.client.load_learnware(learnware_path=zippath, runnable_option="conda_env") for zippath in self.zip_paths] + + learnware_list = [ + self.client.load_learnware(learnware_path=zippath, runnable_option="conda_env") + for zippath in self.zip_paths + ] reuser = AveragingReuser(learnware_list, mode="vote_by_label") input_array = np.random.random(size=(20, 13)) print(reuser.predict(input_array)) @@ -39,7 +41,7 @@ class TestLearnwareLoad(unittest.TestCase): def test_load_multi_learnware_by_zippath(self): for (learnware_id, zip_path) in zip(self.learnware_ids, self.zip_paths): self.client.download_learnware(learnware_id, zip_path) - + learnware_list = self.client.load_learnware(learnware_path=self.zip_paths, runnable_option="conda_env") reuser = AveragingReuser(learnware_list, mode="vote_by_label") input_array = np.random.random(size=(20, 13)) @@ -47,9 +49,11 @@ class TestLearnwareLoad(unittest.TestCase): for learnware in learnware_list: print(learnware.id, learnware.predict(input_array)) - + def test_load_single_learnware_by_id(self): - learnware_list = [self.client.load_learnware(learnware_id=idx, runnable_option="conda_env") for idx in self.learnware_ids] + learnware_list = [ + self.client.load_learnware(learnware_id=idx, runnable_option="conda_env") for idx in self.learnware_ids + ] reuser = AveragingReuser(learnware_list, mode="vote_by_label") input_array = np.random.random(size=(20, 13)) print(reuser.predict(input_array)) @@ -57,7 +61,7 @@ class TestLearnwareLoad(unittest.TestCase): for learnware in learnware_list: print(learnware.id, learnware.predict(input_array)) - def test_load_multi_learnware_by_id(self): + def test_load_multi_learnware_by_id(self): learnware_list = self.client.load_learnware(learnware_id=self.learnware_ids, runnable_option="conda_env") reuser = AveragingReuser(learnware_list, mode="vote_by_label") input_array = np.random.random(size=(20, 13)) @@ -67,5 +71,5 @@ class TestLearnwareLoad(unittest.TestCase): print(learnware.id, learnware.predict(input_array)) -if __name__ == '__main__': - unittest.main() \ No newline at end of file +if __name__ == "__main__": + unittest.main()