Browse Source

[MNT] modify client load_learnware

tags/v0.3.2
Gene 2 years ago
parent
commit
eda296f891
2 changed files with 11 additions and 17 deletions
  1. +5
    -10
      learnware/client/learnware_client.py
  2. +6
    -7
      tests/test_learnware_client/test_load_conda.py

+ 5
- 10
learnware/client/learnware_client.py View File

@@ -287,8 +287,8 @@ class LearnwareClient:
learnware id or learnware id list
runnable_option : str
the option for instantiating learnwares
- "normal": instantiate learnware without installing environment
- "conda_env": instantiate learnware with installing conda virtual environment
- None: instantiate learnware without installing environment
- "conda": instantiate learnware with installing conda virtual environment
- "docker": instantiate learnware with creating docker container

Returns
@@ -296,10 +296,8 @@ class LearnwareClient:
Learnware
The contructed learnware object or object list
"""
if runnable_option is not None and runnable_option not in ["normal", "conda_env", "docker"]:
raise logger.warning(
f"runnable_option must be one of ['normal', 'conda_env', 'docker'], but got {runnable_option}"
)
if runnable_option is not None and runnable_option not in ["conda", "docker"]:
raise logger.warning(f"runnable_option must be one of ['conda', 'docker'], but got {runnable_option}")

if learnware_path is None and learnware_id is None:
raise ValueError("Requires one of learnware_path or learnware_id")
@@ -357,10 +355,7 @@ class LearnwareClient:
learnware_list.append(learnware_obj)

if runnable_option is not None:
if runnable_option == "normal":
for i in range(len(learnware_list)):
learnware_list[i].instantiate_model()
elif runnable_option == "conda_env":
if runnable_option == "conda":
with LearnwaresContainer(learnware_list, cleanup=False, mode="conda") as env_container:
learnware_list = env_container.get_learnwares_with_container()
elif runnable_option == "docker":


+ 6
- 7
tests/test_learnware_client/test_load_conda.py View File

@@ -28,8 +28,7 @@ class TestLearnwareLoad(unittest.TestCase):
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
self.client.load_learnware(learnware_path=zippath, runnable_option="conda") for zippath in self.zip_paths
]
reuser = AveragingReuser(learnware_list, mode="vote_by_label")
input_array = np.random.random(size=(20, 13))
@@ -42,7 +41,7 @@ class TestLearnwareLoad(unittest.TestCase):
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")
learnware_list = self.client.load_learnware(learnware_path=self.zip_paths, runnable_option="conda")
reuser = AveragingReuser(learnware_list, mode="vote_by_label")
input_array = np.random.random(size=(20, 13))
print(reuser.predict(input_array))
@@ -52,7 +51,7 @@ class TestLearnwareLoad(unittest.TestCase):

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
self.client.load_learnware(learnware_id=idx, runnable_option="conda") for idx in self.learnware_ids
]
reuser = AveragingReuser(learnware_list, mode="vote_by_label")
input_array = np.random.random(size=(20, 13))
@@ -62,7 +61,7 @@ class TestLearnwareLoad(unittest.TestCase):
print(learnware.id, learnware.predict(input_array))

def test_load_multi_learnware_by_id(self):
learnware_list = self.client.load_learnware(learnware_id=self.learnware_ids, runnable_option="conda_env")
learnware_list = self.client.load_learnware(learnware_id=self.learnware_ids, runnable_option="conda")
reuser = AveragingReuser(learnware_list, mode="vote_by_label")
input_array = np.random.random(size=(20, 13))
print(reuser.predict(input_array))
@@ -72,13 +71,13 @@ class TestLearnwareLoad(unittest.TestCase):

def test_load_single_learnware_by_id_pip(self):
learnware_id = "00000147"
learnware = self.client.load_learnware(learnware_id=learnware_id, runnable_option="conda_env")
learnware = self.client.load_learnware(learnware_id=learnware_id, runnable_option="conda")
input_array = np.random.random(size=(20, 23))
print(learnware.predict(input_array))

def test_load_single_learnware_by_id_conda(self):
learnware_id = "00000148"
learnware = self.client.load_learnware(learnware_id=learnware_id, runnable_option="conda_env")
learnware = self.client.load_learnware(learnware_id=learnware_id, runnable_option="conda")
input_array = np.random.random(size=(20, 204))
print(learnware.predict(input_array))



Loading…
Cancel
Save