Browse Source

[FIX] fix bugs in LearnwareClient

tags/v0.3.2
Gene 2 years ago
parent
commit
3e1d1997c0
3 changed files with 47 additions and 15 deletions
  1. +1
    -1
      learnware/client/container.py
  2. +1
    -4
      learnware/client/learnware_client.py
  3. +45
    -10
      tests/test_client/test_download.py

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

@@ -31,7 +31,7 @@ class ModelEnvContainer(BaseModel):

with open(model_path, "wb") as model_fp:
pickle.dump(self.model_config, model_fp)
system_execute(
[
"conda",


+ 1
- 4
learnware/client/learnware_client.py View File

@@ -337,11 +337,8 @@ class LearnwareClient:

if load_model:
learnware_obj.instantiate_model()
pass

return learnware_obj
pass
pass

def system(self, command):
retcd = os.system(command)


+ 45
- 10
tests/test_client/test_download.py View File

@@ -1,12 +1,55 @@
import os
import zipfile
import numpy as np

import learnware
from learnware.learnware import get_learnware_from_dirpath
from learnware.client import LearnwareClient
from learnware.client.container import ModelEnvContainer, LearnwaresContainer
from learnware.learnware.reuse import AveragingReuser


def test_container(zip_paths):
semantic_specification = dict()
semantic_specification["Data"] = {"Type": "Class", "Values": ["Text"]}
semantic_specification["Task"] = {"Type": "Class", "Values": ["Ranking"]}
semantic_specification["Library"] = {"Type": "Class", "Values": ["Scikit-learn"]}
semantic_specification["Scenario"] = {"Type": "Tag", "Values": "Financial"}
semantic_specification["Name"] = {"Type": "String", "Values": "test"}
semantic_specification["Description"] = {"Type": "String", "Values": "test"}
learnware_list = []
for id, zip_path in enumerate(zip_paths):
dir_path = zip_path[:-4]
with zipfile.ZipFile(zip_path, "r") as z_file:
z_file.extractall(dir_path)

learnware = get_learnware_from_dirpath(f"test_id{id}", semantic_specification, dir_path)
learnware_list.append(learnware)
with LearnwaresContainer(learnware_list, zip_paths) as env_container:
learnware_list = env_container.get_learnware_list_with_container()
reuser = AveragingReuser(learnware_list, mode="vote_by_label")
input_array = np.random.random(size=(20, 13))
print(reuser.predict(input_array))
for idx, learnware in enumerate(learnware_list):
print(f"learnware_{idx}", learnware.predict(input_array))

def test_load(zip_paths):
learnware_list = [client.load_learnware(file, load_model=False) for file in zip_paths]
with LearnwaresContainer(learnware_list, zip_paths) as env_container:
learnware_list = env_container.get_learnware_list_with_container()
reuser = AveragingReuser(learnware_list, mode="vote_by_label")
input_array = np.random.random(size=(20, 13))
print(reuser.predict(input_array))
for idx, learnware in enumerate(learnware_list):
print(f"learnware_{idx}", learnware.predict(input_array))

if __name__ == "__main__":
email = "liujd@lamda.nju.edu.cn"
token = "f7e647146a314c6e8b4e2e1079c4bca4"
@@ -21,13 +64,5 @@ if __name__ == "__main__":
zip_paths[i] = os.path.join(root, zip_paths[i])
client.download_learnware(learnware_ids[i], zip_paths[i])
learnware_list = [client.load_learnware(file, load_model=False) for file in zip_paths]

with LearnwaresContainer(learnware_list, zip_paths) as env_container:
learnware_list = env_container.get_learnware_list_with_container()
reuser = AveragingReuser(learnware_list, mode="vote_by_label")
input_array = np.random.random(size=(20, 13))
print(reuser.predict(input_array))
for idx, learnware in enumerate(learnware_list):
print(f"learnware_{idx}", reuser.predict(learnware))
test_container(zip_paths)
# test_load(zip_paths)

Loading…
Cancel
Save