Browse Source

[FIX] fix cannot import generate_rkme bugs

tags/v0.3.2
bxdd 2 years ago
parent
commit
c31f5f4dc8
9 changed files with 25 additions and 35 deletions
  1. +2
    -2
      docs/workflow/submit.rst
  2. +3
    -4
      examples/dataset_m5_workflow/main.py
  3. +3
    -5
      examples/dataset_pfs_workflow/main.py
  4. +4
    -9
      examples/dataset_text_workflow/main.py
  5. +3
    -4
      examples/workflow_by_code/main.py
  6. +1
    -1
      learnware/reuse/job_selector.py
  7. +3
    -3
      tests/test_market/test_easy.py
  8. +2
    -3
      tests/test_specification/test_rkme.py
  9. +4
    -4
      tests/test_workflow/test_workflow.py

+ 2
- 2
docs/workflow/submit.rst View File

@@ -80,10 +80,10 @@ the following code snippet offers guidance on how to construct and store the RKM

.. code-block:: python
import learnware.specification as specification
from learnware.specification import generate_rkme_spec
# generate rkme specification for digits dataset
spec = specification.utils.generate_rkme_spec(X=data_X)
spec = generate_rkme_spec(X=data_X)
spec.save("stat.json")

Significantly, the RKME generation process is entirely conducted on your local machine, without any involvement of cloud services,


+ 3
- 4
examples/dataset_m5_workflow/main.py View File

@@ -9,9 +9,8 @@ from shutil import copyfile, rmtree
import learnware
from learnware.market import EasyMarket, BaseUserInfo
from learnware.market import database_ops
from learnware.learnware import Learnware
from learnware.reuse import JobSelectorReuser, AveragingReuser
import learnware.specification as specification
from learnware.specification import generate_rkme_spec
from m5 import DataLoader
from learnware.logger import get_module_logger

@@ -88,7 +87,7 @@ class M5DatasetWorkflow:
for idx in tqdm(idx_list):
train_x, train_y, test_x, test_y = m5.get_idx_data(idx)
st = time.time()
spec = specification.utils.generate_rkme_spec(X=train_x, gamma=0.1, cuda_idx=0)
spec = generate_rkme_spec(X=train_x, gamma=0.1, cuda_idx=0)
ed = time.time()
logger.info("Stat spec generated in %.3f s" % (ed - st))

@@ -140,7 +139,7 @@ class M5DatasetWorkflow:

for idx in idx_list:
train_x, train_y, test_x, test_y = m5.get_idx_data(idx)
user_spec = specification.utils.generate_rkme_spec(X=test_x, gamma=0.1, cuda_idx=0)
user_spec = generate_rkme_spec(X=test_x, gamma=0.1, cuda_idx=0)
user_spec_path = f"./user_spec/user_{idx}.json"
user_spec.save(user_spec_path)



+ 3
- 5
examples/dataset_pfs_workflow/main.py View File

@@ -8,10 +8,8 @@ from shutil import copyfile, rmtree

import learnware
from learnware.market import EasyMarket, BaseUserInfo
from learnware.market import database_ops
from learnware.learnware import Learnware
from learnware.reuse import JobSelectorReuser, AveragingReuser
import learnware.specification as specification
from learnware.specification import generate_rkme_spec
from pfs import Dataloader
from learnware.logger import get_module_logger

@@ -86,7 +84,7 @@ class PFSDatasetWorkflow:
for idx in tqdm(idx_list):
train_x, train_y, test_x, test_y = pfs.get_idx_data(idx)
st = time.time()
spec = specification.utils.generate_rkme_spec(X=train_x, gamma=0.1, cuda_idx=0)
spec = generate_rkme_spec(X=train_x, gamma=0.1, cuda_idx=0)
ed = time.time()
logger.info("Stat spec generated in %.3f s" % (ed - st))

@@ -138,7 +136,7 @@ class PFSDatasetWorkflow:

for idx in idx_list:
train_x, train_y, test_x, test_y = pfs.get_idx_data(idx)
user_spec = specification.utils.generate_rkme_spec(X=test_x, gamma=0.1, cuda_idx=0)
user_spec = generate_rkme_spec(X=test_x, gamma=0.1, cuda_idx=0)
user_spec_path = f"./user_spec/user_{idx}.json"
user_spec.save(user_spec_path)



+ 4
- 9
examples/dataset_text_workflow/main.py View File

@@ -10,9 +10,7 @@ import time
import pickle

from learnware.market import instantiate_learnware_market, BaseUserInfo
from learnware.market import database_ops
from learnware.learnware import Learnware
import learnware.specification as specification
from learnware.specification import RKMETextSpecification
from learnware.logger import get_module_logger

from shutil import copyfile, rmtree
@@ -99,8 +97,7 @@ def prepare_learnware(data_path, model_path, init_file_path, yaml_path, save_roo
semantic_spec = semantic_specs[0]

st = time.time()
# user_spec = specification.utils.generate_rkme_spec(X=X, gamma=0.1, cuda_idx=0)
user_spec = specification.RKMETextSpecification()
user_spec = RKMETextSpecification()
user_spec.generate_stat_spec_from_data(X=X)
ed = time.time()
logger.info("Stat spec generated in %.3f s" % (ed - st))
@@ -163,10 +160,8 @@ def test_search(gamma=0.1, load_market=True):
user_data = pickle.load(f)
with open(user_label_path, "rb") as f:
user_label = pickle.load(f)
# user_data = np.load(user_data_path)
# user_label = np.load(user_label_path)
# user_stat_spec = specification.utils.generate_rkme_spec(X=user_data, gamma=gamma, cuda_idx=0)
user_stat_spec = specification.RKMETextSpecification()

user_stat_spec = RKMETextSpecification()
user_stat_spec.generate_stat_spec_from_data(X=user_data)
user_info = BaseUserInfo(semantic_spec=user_semantic, stat_info={"RKMETextSpecification": user_stat_spec})
logger.info("Searching Market for user: %d" % (i))


+ 3
- 4
examples/workflow_by_code/main.py View File

@@ -12,8 +12,7 @@ from shutil import copyfile, rmtree
import learnware
from learnware.market import EasyMarket, BaseUserInfo
from learnware.reuse import JobSelectorReuser, AveragingReuser
import learnware.specification as specification
from learnware.utils import get_module_by_module_path
from learnware.specification import generate_rkme_spec

curr_root = os.path.dirname(os.path.abspath(__file__))

@@ -54,7 +53,7 @@ class LearnwareMarketWorkflow:

joblib.dump(clf, os.path.join(dir_path, "svm.pkl"))

spec = specification.utils.generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
spec = generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
spec.save(os.path.join(dir_path, "svm.json"))

init_file = os.path.join(dir_path, "__init__.py")
@@ -174,7 +173,7 @@ class LearnwareMarketWorkflow:
X, y = load_digits(return_X_y=True)
_, data_X, _, data_y = train_test_split(X, y, test_size=0.3, shuffle=True)

stat_spec = specification.utils.generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
stat_spec = generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
user_info = BaseUserInfo(semantic_spec=user_semantic, stat_info={"RKMETableSpecification": stat_spec})

_, _, _, mixture_learnware_list = easy_market.search_learnware(user_info)


+ 1
- 1
learnware/reuse/job_selector.py View File

@@ -11,7 +11,7 @@ from .base import BaseReuser
from ..market.utils import parse_specification_type
from ..learnware import Learnware
from ..specification import RKMETableSpecification, RKMETextSpecification
from ..specification.utils import generate_rkme_spec
from ..specification import generate_rkme_spec
from ..logger import get_module_logger

logger = get_module_logger("job_selector_reuse")


+ 3
- 3
tests/test_market/test_easy.py View File

@@ -12,7 +12,7 @@ from shutil import copyfile, rmtree

import learnware
from learnware.market import instantiate_learnware_market, BaseUserInfo
import learnware.specification as specification
from learnware.specification import RKMETableSpecification, generate_rkme_spec

curr_root = os.path.dirname(os.path.abspath(__file__))

@@ -62,7 +62,7 @@ class TestMarket(unittest.TestCase):

joblib.dump(clf, os.path.join(dir_path, "svm.pkl"))

spec = specification.utils.generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
spec = generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
spec.save(os.path.join(dir_path, "svm.json"))

init_file = os.path.join(dir_path, "__init__.py")
@@ -170,7 +170,7 @@ class TestMarket(unittest.TestCase):
with zipfile.ZipFile(zip_path, "r") as zip_obj:
zip_obj.extractall(path=unzip_dir)

user_spec = specification.rkme.RKMETableSpecification()
user_spec = RKMETableSpecification()
user_spec.load(os.path.join(unzip_dir, "svm.json"))
user_info = BaseUserInfo(semantic_spec=user_semantic, stat_info={"RKMETableSpecification": user_spec})
(


+ 2
- 3
tests/test_specification/test_rkme.py View File

@@ -7,9 +7,8 @@ import unittest
import tempfile
import numpy as np

import learnware.specification as specification
from learnware.specification import RKMETableSpecification, RKMEImageSpecification, RKMETextSpecification
from learnware.specification import generate_rkme_image_spec, generate_rkme_spec
from learnware.specification import generate_rkme_image_spec, generate_rkme_spec, generate_rkme_text_spec


class TestRKME(unittest.TestCase):
@@ -71,7 +70,7 @@ class TestRKME(unittest.TestCase):
return text_list

def _test_text_rkme(X):
rkme = specification.utils.generate_rkme_text_spec(X)
rkme = generate_rkme_text_spec(X)

with tempfile.TemporaryDirectory(prefix="learnware_") as tempdir:
rkme_path = os.path.join(tempdir, "rkme.json")


+ 4
- 4
tests/test_workflow/test_workflow.py View File

@@ -13,7 +13,7 @@ from shutil import copyfile, rmtree
import learnware
from learnware.market import EasyMarket, BaseUserInfo
from learnware.reuse import JobSelectorReuser, AveragingReuser, EnsemblePruningReuser
import learnware.specification as specification
from learnware.specification import generate_rkme_spec, RKMETableSpecification

curr_root = os.path.dirname(os.path.abspath(__file__))

@@ -57,7 +57,7 @@ class TestAllWorkflow(unittest.TestCase):

joblib.dump(clf, os.path.join(dir_path, "svm.pkl"))

spec = specification.utils.generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
spec = generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
spec.save(os.path.join(dir_path, "svm.json"))

init_file = os.path.join(dir_path, "__init__.py")
@@ -159,7 +159,7 @@ class TestAllWorkflow(unittest.TestCase):
with zipfile.ZipFile(zip_path, "r") as zip_obj:
zip_obj.extractall(path=unzip_dir)

user_spec = specification.RKMETableSpecification()
user_spec = RKMETableSpecification()
user_spec.load(os.path.join(unzip_dir, "svm.json"))
user_info = BaseUserInfo(semantic_spec=user_semantic, stat_info={"RKMETableSpecification": user_spec})
(
@@ -185,7 +185,7 @@ class TestAllWorkflow(unittest.TestCase):
X, y = load_digits(return_X_y=True)
train_X, data_X, train_y, data_y = train_test_split(X, y, test_size=0.3, shuffle=True)

stat_spec = specification.utils.generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
stat_spec = generate_rkme_spec(X=data_X, gamma=0.1, cuda_idx=0)
user_info = BaseUserInfo(semantic_spec=user_semantic, stat_info={"RKMETableSpecification": stat_spec})

_, _, _, mixture_learnware_list = easy_market.search_learnware(user_info)


Loading…
Cancel
Save