From 661796fb2bb0af091a02f6e36088a3e70a1e168d Mon Sep 17 00:00:00 2001 From: bxdd Date: Sun, 15 Oct 2023 17:09:45 +0800 Subject: [PATCH] [ENH | MNT] add cleanup option, and make model auto cleanup when exit --- learnware/client/container.py | 7 +++++-- tests/test_learnware_client/test_learnware.py | 7 +++---- tests/test_learnware_client/test_reuse.py | 14 +++++++------- 3 files changed, 15 insertions(+), 13 deletions(-) diff --git a/learnware/client/container.py b/learnware/client/container.py index 16ecf49..8b03aee 100644 --- a/learnware/client/container.py +++ b/learnware/client/container.py @@ -1,6 +1,5 @@ import os import pickle -import atexit import tempfile import shortuuid from concurrent.futures import ProcessPoolExecutor @@ -110,7 +109,7 @@ class ModelEnvContainer(BaseModel): class LearnwaresContainer: - def __init__(self, learnwares: Union[List[Learnware], Learnware], learnware_zippaths: Union[List[str], str]): + def __init__(self, learnwares: Union[List[Learnware], Learnware], learnware_zippaths: Union[List[str], str], cleanup=True): """The initializaiton method for base reuser Parameters @@ -132,6 +131,7 @@ class LearnwaresContainer: ) for _learnware, _zippath in zip(learnwares, learnware_zippaths) ] + self.cleanup = cleanup def __enter__(self): model_list = [_learnware.get_model() for _learnware in self.learnware_list] @@ -139,6 +139,9 @@ class LearnwaresContainer: executor.map(self._initialize_model_container, model_list) def __exit__(self, exc_type, exc_val, exc_tb): + if not self.cleanup: + logger.warning(f"Notice, the learnware container env is not clean up!") + return 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._destroy_model_container, model_list) diff --git a/tests/test_learnware_client/test_learnware.py b/tests/test_learnware_client/test_learnware.py index 73a7d53..f403ef4 100644 --- a/tests/test_learnware_client/test_learnware.py +++ b/tests/test_learnware_client/test_learnware.py @@ -15,7 +15,6 @@ if __name__ == "__main__": z_file.extractall(learnware_dirpath) learnware = get_learnware_from_dirpath(id='test', semantic_spec=semantic_specification, learnware_dirpath=learnware_dirpath) - env_container = LearnwaresContainer(learnware, zip_path) - learnware = env_container.get_learnwares_with_container()[0] - - EasyMarket.check_learnware(learnware) + with LearnwaresContainer(learnware, zip_path) as env_container: + learnware = env_container.get_learnwares_with_container()[0] + EasyMarket.check_learnware(learnware) diff --git a/tests/test_learnware_client/test_reuse.py b/tests/test_learnware_client/test_reuse.py index 9fdf003..d0e4d9a 100644 --- a/tests/test_learnware_client/test_reuse.py +++ b/tests/test_learnware_client/test_reuse.py @@ -25,10 +25,10 @@ if __name__ == "__main__": learnware = get_learnware_from_dirpath(f"test_id{id}", semantic_specification, dir_path) learnware_list.append(learnware) - env_container = LearnwaresContainer(learnware_list, zip_paths) - learnware_list = env_container.get_learnwares_with_container() - reuser = AveragingReuser(learnware_list, mode="vote") - input_array = np.random.randint(0, 3, size=(20, 9)) - print(reuser.predict(input_array).argmax(axis=1)) - for id, ind_learner in enumerate(learnware_list): - print(f"learner_{id}", reuser.predict(input_array).argmax(axis=1)) + with LearnwaresContainer(learnware_list, zip_paths) as env_container: + learnware_list = env_container.get_learnwares_with_container() + reuser = AveragingReuser(learnware_list, mode="vote") + input_array = np.random.randint(0, 3, size=(20, 9)) + print(reuser.predict(input_array).argmax(axis=1)) + for id, ind_learner in enumerate(learnware_list): + print(f"learner_{id}", reuser.predict(input_array).argmax(axis=1))