Browse Source

[MNT] format code

tags/v0.3.2
Gene 2 years ago
parent
commit
e062a181cf
3 changed files with 30 additions and 21 deletions
  1. +4
    -4
      learnware/client/container.py
  2. +13
    -8
      learnware/client/learnware_client.py
  3. +13
    -9
      tests/test_client/test_load.py

+ 4
- 4
learnware/client/container.py View File

@@ -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())
self._destroy_model_container(_learnware.get_model())

+ 13
- 8
learnware/client/learnware_client.py View File

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


+ 13
- 9
tests/test_client/test_load.py View File

@@ -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()
if __name__ == "__main__":
unittest.main()

Loading…
Cancel
Save