Browse Source

[ENH | MNT] add cleanup option, and make model auto cleanup when exit

tags/v0.3.2
bxdd 2 years ago
parent
commit
661796fb2b
3 changed files with 15 additions and 13 deletions
  1. +5
    -2
      learnware/client/container.py
  2. +3
    -4
      tests/test_learnware_client/test_learnware.py
  3. +7
    -7
      tests/test_learnware_client/test_reuse.py

+ 5
- 2
learnware/client/container.py View File

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


+ 3
- 4
tests/test_learnware_client/test_learnware.py View File

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

+ 7
- 7
tests/test_learnware_client/test_reuse.py View File

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

Loading…
Cancel
Save