From 568ef0f4369c485a28ad29f994521485861e8b92 Mon Sep 17 00:00:00 2001 From: bxdd Date: Mon, 23 Oct 2023 22:07:17 +0800 Subject: [PATCH 01/35] [MNT] refactor market: for temp save --- learnware/market/base.py | 91 +++++++++++++++++++++++++++++++++++++++- 1 file changed, 90 insertions(+), 1 deletion(-) diff --git a/learnware/market/base.py b/learnware/market/base.py index 75f6228..e03e3c4 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -64,7 +64,8 @@ class BaseMarket: raise NotImplementedError("reload market is Not Implemented") - def check_learnware(self, learnware: Learnware) -> bool: + @classmethod + def check_learnware(cls, learnware: Learnware) -> bool: """Check the utility of a learnware Parameters @@ -195,3 +196,91 @@ class BaseMarket: """ raise NotImplementedError("get semantic spec list is not implemented") + + +class LearnwareOrganizer: + def __init__(self, market_id): + self.market_id = market_id + + + def reload_market(self) -> bool: + """Reload the market when server restared. + + Parameters + ---------- + market_path : str + Directory for market data. '_IP_:_port_' for loading from database. + semantic_spec_list_path : str + Directory for available semantic_spec. Should be a json file. + + Returns + ------- + bool + A flag indicating whether the market is reload successfully. + """ + + raise NotImplementedError("reload market is Not Implemented") + + def add_learnware( + self, learnware_name: str, model_path: str, stat_spec_path: str, semantic_spec: dict, desc: str + ) -> Tuple[str, bool]: + """Add a learnware into the market. + + .. note:: + + Given a prediction of a certain time, all signals before this time will be prepared well. + + + Parameters + ---------- + learnware_name : str + Name of new learnware. + model_path : str + Filepath for learnware model, a zipped file. + stat_spec_path : str + Filepath for statistical specification, a '.npy' file. + How to pass parameters requires further discussion. + semantic_spec : dict + semantic_spec for new learnware, in dictionary format. + desc : str + Brief desciption for new learnware. + + Returns + ------- + Tuple[str, bool] + str indicating model_id, bool indicating whether the learnware is added successfully. + + Raises + ------ + FileNotFoundError + file for model or statistical specification not found + + """ + raise NotImplementedError("add learnware is Not Implemented") + + + def delete_learnware(self, id: str) -> bool: + """Delete a learnware from market + + Parameters + ---------- + id : str + id of learnware to be deleted + + Returns + ------- + bool + True if the target learnware is deleted successfully. + + Raises + ------ + Exception + Raise an excpetion when given id is NOT found in learnware list + """ + raise NotImplementedError("delete learnware is Not Implemented") + +class LearnwareSearcher: + def __init__(self, learnware_organizor): + + def search_learnware(self, user_info: BaseUserInfo) -> Tuple[Any, List[Learnware]]: + pass \ No newline at end of file From 4e4912cf5a9f66de68a598366df8166cf250f372 Mon Sep 17 00:00:00 2001 From: bxdd Date: Tue, 24 Oct 2023 22:49:08 +0800 Subject: [PATCH 02/35] [ENH] add organizer, searcher, checker for market --- learnware/market/base.py | 79 ++++++++++++++++++------------ learnware/market/easy/__init__.py | 7 +++ learnware/market/easy/checker.py | 72 +++++++++++++++++++++++++++ learnware/market/easy/organizer.py | 11 +++++ learnware/market/easy/searcher.py | 8 +++ learnware/market/searcher.py | 10 ++++ 6 files changed, 155 insertions(+), 32 deletions(-) create mode 100644 learnware/market/easy/__init__.py create mode 100644 learnware/market/easy/checker.py create mode 100644 learnware/market/easy/organizer.py create mode 100644 learnware/market/easy/searcher.py create mode 100644 learnware/market/searcher.py diff --git a/learnware/market/base.py b/learnware/market/base.py index e03e3c4..3998612 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -1,11 +1,14 @@ import os +import torch +import traceback import numpy as np -import pandas as pd -from typing import Tuple, Any, List, Union, Dict + +from typing import Tuple, Any, List, Union from ..learnware import Learnware -from ..specification import RKMEStatSpecification +from ..logger import get_module_logger +logger = get_module_logger("market_base", "INFO") class BaseUserInfo: """User Information for searching learnware""" @@ -43,19 +46,13 @@ class BaseUserInfo: class BaseMarket: """Base interface for market, it provide the interface of search/add/detele/update learnwares""" - def __init__(self, market_id: str = None): + def __init__(self, market_id: str = None, checker: 'LearnwareChecker' = None): self.market_id = market_id + + self.learnware_checker = LearnwareChecker() if checker is None else checker - def reload_market(self, market_path: str, semantic_spec_list_path: str) -> bool: + def reload_market(self, **kwargs) -> bool: """Reload the market when server restared. - - Parameters - ---------- - market_path : str - Directory for market data. '_IP_:_port_' for loading from database. - semantic_spec_list_path : str - Directory for available semantic_spec. Should be a json file. - Returns ------- bool @@ -64,8 +61,7 @@ class BaseMarket: raise NotImplementedError("reload market is Not Implemented") - @classmethod - def check_learnware(cls, learnware: Learnware) -> bool: + def check_learnware(self, learnware: Learnware) -> bool: """Check the utility of a learnware Parameters @@ -77,7 +73,7 @@ class BaseMarket: bool A flag indicating whether the learnware can be accepted. """ - return True + return self.learnware_checker(learnware) def add_learnware( self, learnware_name: str, model_path: str, stat_spec_path: str, semantic_spec: dict, desc: str @@ -221,9 +217,7 @@ class LearnwareOrganizer: raise NotImplementedError("reload market is Not Implemented") - def add_learnware( - self, learnware_name: str, model_path: str, stat_spec_path: str, semantic_spec: dict, desc: str - ) -> Tuple[str, bool]: + def add_learnware(self, zip_path: str, semantic_spec: dict) -> Tuple[str, bool]: """Add a learnware into the market. .. note:: @@ -233,22 +227,17 @@ class LearnwareOrganizer: Parameters ---------- - learnware_name : str - Name of new learnware. - model_path : str + zip_path : str Filepath for learnware model, a zipped file. - stat_spec_path : str - Filepath for statistical specification, a '.npy' file. - How to pass parameters requires further discussion. semantic_spec : dict semantic_spec for new learnware, in dictionary format. - desc : str - Brief desciption for new learnware. Returns ------- - Tuple[str, bool] - str indicating model_id, bool indicating whether the learnware is added successfully. + Tuple[str, int] + - str indicating model_id + - int indicating what the flag of learnware is added. + Raises ------ @@ -280,7 +269,33 @@ class LearnwareOrganizer: raise NotImplementedError("delete learnware is Not Implemented") class LearnwareSearcher: - def __init__(self, learnware_organizor): + def __init__(self, organizer): + self.learnware_organizer = organizer + + def __call__(self, user_info: BaseUserInfo): + raise NotImplementedError("'__call__' method is not implemented in LearnwareSearcher") - def search_learnware(self, user_info: BaseUserInfo) -> Tuple[Any, List[Learnware]]: - pass \ No newline at end of file + +class LearnwareChecker: + INVALID_LEARNWARE = -1 + NONUSABLE_LEARNWARE = 0 + USABLE_LEARWARE = 1 + + @classmethod + def __call__(cls, learnware: Learnware) -> int: + """Check the utility of a learnware + + Parameters + ---------- + learnware : Learnware + + Returns + ------- + int + A flag indicating whether the learnware can be accepted. + - The INVALID_LEARNWARE denotes the learnware does not pass the check + - The NOPREDICTION_LEARNWARE denotes the learnware pass the check but cannot make prediction due to some env dependency + - The NOPREDICTION_LEARNWARE denotes the leanrware pass the check and can make prediction + """ + + raise NotImplementedError("'__call__' method is not implemented in LearnwareChecker") \ No newline at end of file diff --git a/learnware/market/easy/__init__.py b/learnware/market/easy/__init__.py new file mode 100644 index 0000000..96f0a34 --- /dev/null +++ b/learnware/market/easy/__init__.py @@ -0,0 +1,7 @@ +from ..base import LearnwareSearcher, LearnwareOrganizer + +class EasySearcher(LearnwareSearcher): + pass + +class EasyOrganizer(LearnwareOrganizer): + pass \ No newline at end of file diff --git a/learnware/market/easy/checker.py b/learnware/market/easy/checker.py new file mode 100644 index 0000000..88d4266 --- /dev/null +++ b/learnware/market/easy/checker.py @@ -0,0 +1,72 @@ +import traceback + +from ..base import LearnwareChecker +from ...logger import get_module_logger + +logger = get_module_logger("easy_checker", "INFO") + +class EasyChecker(LearnwareChecker): + + @classmethod + def __call__(cls, learnware): + semantic_spec = learnware.get_specification().get_semantic_spec() + + try: + # check model instantiation + learnware.instantiate_model() + + except Exception as e: + traceback.print_exc() + logger.warning(f"The learnware [{learnware.id}] is instantiated failed! Due to {e}") + return cls.NONUSABLE_LEARNWARE + + try: + learnware_model = learnware.get_model() + + # check input shape + if semantic_spec["Data"]["Values"][0] == "Table": + input_shape = (semantic_spec["Input"]["Dimension"],) + else: + input_shape = learnware_model.input_shape + pass + + # check rkme dimension + stat_spec = learnware.get_specification().get_stat_spec_by_name("RKMEStatSpecification") + if stat_spec is not None: + if stat_spec.get_z().shape[1:] != input_shape: + logger.warning(f"The learnware [{learnware.id}] input dimension mismatch with stat specification") + return cls.NONUSABLE_LEARNWARE + pass + + inputs = np.random.randn(10, *input_shape) + outputs = learnware.predict(inputs) + + # check output + if outputs.ndim == 1: + outputs = outputs.reshape(-1, 1) + pass + + if semantic_spec["Task"]["Values"][0] in ("Classification", "Regression", "Feature Extraction"): + # check output type + if isinstance(outputs, torch.Tensor): + outputs = outputs.detach().cpu().numpy() + if not isinstance(outputs, np.ndarray): + logger.warning(f"The learnware [{learnware.id}] output must be np.ndarray or torch.Tensor") + return cls.NONUSABLE_LEARNWARE + + # check output shape + output_dim = int(semantic_spec["Output"]["Dimension"]) + if outputs[0].shape[0] != output_dim: + logger.warning(f"The learnware [{learnware.id}] input and output dimention is error") + return cls.NONUSABLE_LEARNWARE + pass + else: + if outputs.shape[1:] != learnware_model.output_shape: + logger.warning(f"The learnware [{learnware.id}] input and output dimention is error") + return cls.NONUSABLE_LEARNWARE + + except Exception as e: + logger.warning(f"The learnware [{learnware.id}] prediction is not avaliable! Due to {repr(e)}") + return cls.NONUSABLE_LEARNWARE + + return cls.USABLE_LEARWARE diff --git a/learnware/market/easy/organizer.py b/learnware/market/easy/organizer.py new file mode 100644 index 0000000..cdbaf69 --- /dev/null +++ b/learnware/market/easy/organizer.py @@ -0,0 +1,11 @@ +import traceback + +from ..base import LearnwareOrganizer +from ...logger import get_module_logger + +logger = get_module_logger("easy_organizer") + + +class EasyOrganizer(LearnwareOrganizer): + + \ No newline at end of file diff --git a/learnware/market/easy/searcher.py b/learnware/market/easy/searcher.py new file mode 100644 index 0000000..eb97a7f --- /dev/null +++ b/learnware/market/easy/searcher.py @@ -0,0 +1,8 @@ +from ..base import LearnwareSearcher +from ...logger import get_module_logger + +logger = get_module_logger('easy_seacher') + +class EasySearcher(LearnwareSearcher): + pass + \ No newline at end of file diff --git a/learnware/market/searcher.py b/learnware/market/searcher.py new file mode 100644 index 0000000..09e6f30 --- /dev/null +++ b/learnware/market/searcher.py @@ -0,0 +1,10 @@ + + +from typing import Tuple, Any, List + +from .base import BaseUserInfo +from ..learnware import Learnware +from ..logger import get_module_logger + +logger = get_module_logger('model') + From 73d6b9228dc0efae55241259553b7561acd51dd1 Mon Sep 17 00:00:00 2001 From: bxdd Date: Wed, 25 Oct 2023 11:00:37 +0800 Subject: [PATCH 03/35] [MNT] modify market interface --- learnware/market/base.py | 138 ++++++++++++++++++++++--------- learnware/market/easy/checker.py | 20 ++--- 2 files changed, 110 insertions(+), 48 deletions(-) diff --git a/learnware/market/base.py b/learnware/market/base.py index 3998612..33774d5 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -46,21 +46,30 @@ class BaseUserInfo: class BaseMarket: """Base interface for market, it provide the interface of search/add/detele/update learnwares""" - def __init__(self, market_id: str = None, checker: 'LearnwareChecker' = None): + def __init__( + self, + market_id: str = None, + organizer: 'LearnwareOrganizer' = None, + checker: 'LearnwareChecker' = None, + searcher: 'LearnwareSearcher' = None, + ): self.market_id = market_id - + self.learnware_organizer = LearnwareOrganizer() if organizer is None else organizer + self.learnware_organizer.reset(market_id=market_id) self.learnware_checker = LearnwareChecker() if checker is None else checker + self.learnware_checker.reset(organizer=self.learnware_organizer) + self.learnware_searcher = LearnwareSearcher() if searcher is None else searcher + self.learnware_searcher.reset(organizer=self.learnware_organizer) - def reload_market(self, **kwargs) -> bool: + def reload_market(self, *args, **kwargs) -> bool: """Reload the market when server restared. Returns ------- bool A flag indicating whether the market is reload successfully. """ - - raise NotImplementedError("reload market is Not Implemented") - + self.learnware_organizer.reload_market(*args, **kwargs) + def check_learnware(self, learnware: Learnware) -> bool: """Check the utility of a learnware @@ -75,30 +84,20 @@ class BaseMarket: """ return self.learnware_checker(learnware) - def add_learnware( - self, learnware_name: str, model_path: str, stat_spec_path: str, semantic_spec: dict, desc: str - ) -> Tuple[str, bool]: + def add_learnware(self, zip_path: str, semantic_spec: dict, **kwargs) -> Tuple[str, bool]: """Add a learnware into the market. .. note:: Given a prediction of a certain time, all signals before this time will be prepared well. - Parameters ---------- - learnware_name : str - Name of new learnware. - model_path : str + zip_path : str Filepath for learnware model, a zipped file. - stat_spec_path : str - Filepath for statistical specification, a '.npy' file. - How to pass parameters requires further discussion. semantic_spec : dict semantic_spec for new learnware, in dictionary format. - desc : str - Brief desciption for new learnware. - + Returns ------- Tuple[str, bool] @@ -110,7 +109,7 @@ class BaseMarket: file for model or statistical specification not found """ - raise NotImplementedError("add learnware is Not Implemented") + return self.learnware_organizer.add_learnware(zip_path, semantic_spec, **kwargs) def search_learnware(self, user_info: BaseUserInfo) -> Tuple[Any, List[Learnware]]: """Search Learnware based on user_info @@ -129,7 +128,7 @@ class BaseMarket: - second is a list of matched learnwares """ - raise NotImplementedError("search learnware is Not Implemented") + return self.learnware_searcher(user_info=user_info) def get_learnware_by_ids(self, id: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """ @@ -148,7 +147,7 @@ class BaseMarket: - The returned items are search results. - 'None' indicating the target id not found. """ - raise NotImplementedError("search learnware is Not Implemented") + return self.learnware_organizer.get_learnware_by_ids(id) def delete_learnware(self, id: str) -> bool: """Delete a learnware from market @@ -168,9 +167,9 @@ class BaseMarket: Exception Raise an excpetion when given id is NOT found in learnware list """ - raise NotImplementedError("delete learnware is Not Implemented") + return self.learnware_organizer.delete_learnware(id) - def update_learnware(self, id: str) -> bool: + def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, **kwargs) -> bool: """ Update Learnware with id and content to be updated. Empty interface. TODO @@ -180,7 +179,7 @@ class BaseMarket: id : str id of target learnware. """ - raise NotImplementedError("update learnware is Not Implemented") + return self.learnware_organizer.update_learnware(id, zip_path=zip_path, semantic_spec=semantic_spec, **kwargs) def get_semantic_spec_list(self) -> dict: """Return all semantic specifications available @@ -198,17 +197,12 @@ class LearnwareOrganizer: def __init__(self, market_id): self.market_id = market_id + def reset(self, market_id): + self.market_id = market_id def reload_market(self) -> bool: - """Reload the market when server restared. - - Parameters - ---------- - market_path : str - Directory for market data. '_IP_:_port_' for loading from database. - semantic_spec_list_path : str - Directory for available semantic_spec. Should be a json file. - + """Reload the learnware organizer when server restared. + Returns ------- bool @@ -268,10 +262,73 @@ class LearnwareOrganizer: """ raise NotImplementedError("delete learnware is Not Implemented") + def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, **kwargs) -> bool: + """ + Update Learnware with id and content to be updated. + Empty interface. TODO + + Parameters + ---------- + id : str + id of target learnware. + """ + raise NotImplementedError("update learnware is Not Implemented") + + def get_semantic_spec_list(self) -> dict: + """Return all semantic specifications available + + Returns + ------- + dict + All emantic specifications in dictionary format + + """ + raise NotImplementedError("get semantic spec list is not implemented") + + def get_learnware_by_ids(self, id: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: + """ + Get Learnware from market by id + + Parameters + ---------- + id : Union[str, List[str]] + Given one id or a list of ids as target. + + Returns + ------- + Union[Learnware, List[Learnware]] + Return a Learnware object or a list of Learnware objects based on the type of input param. + + - The returned items are search results. + - 'None' indicating the target id not found. + """ + raise NotImplementedError("get_learnware_by_ids is not implemented") + + def get_learnware_path_by_ids(self, ids: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: + """Get Zipped Learnware file by id + + Parameters + ---------- + ids : Union[str, List[str]] + Give a id or a list of ids + str: id of targer learware + List[str]: A list of ids of target learnwares + + Returns + ------- + Union[Learnware, List[Learnware]] + Return the path for target learnware or list of path. + None for Learnware NOT Found. + """ + raise NotImplementedError("get_learnware_path_by_ids is not implemented") + class LearnwareSearcher: - def __init__(self, organizer): - self.learnware_organizer = organizer + def __init__(self, organizer: LearnwareOrganizer = None): + self.learnware_oganizer = organizer + def reset(self, organizer): + self.learnware_oganizer = organizer + def __call__(self, user_info: BaseUserInfo): raise NotImplementedError("'__call__' method is not implemented in LearnwareSearcher") @@ -281,8 +338,13 @@ class LearnwareChecker: NONUSABLE_LEARNWARE = 0 USABLE_LEARWARE = 1 - @classmethod - def __call__(cls, learnware: Learnware) -> int: + def __init__(self, organizer: LearnwareOrganizer = None): + self.learnware_oganizer = organizer + + def reset(self, organizer): + self.learnware_oganizer = organizer + + def __call__(self, learnware: Learnware) -> int: """Check the utility of a learnware Parameters diff --git a/learnware/market/easy/checker.py b/learnware/market/easy/checker.py index 88d4266..50fac94 100644 --- a/learnware/market/easy/checker.py +++ b/learnware/market/easy/checker.py @@ -1,14 +1,14 @@ import traceback -from ..base import LearnwareChecker +from ..base import LearnwareChecker, LearnwareOrganizer from ...logger import get_module_logger logger = get_module_logger("easy_checker", "INFO") class EasyChecker(LearnwareChecker): + - @classmethod - def __call__(cls, learnware): + def __call__(self, learnware): semantic_spec = learnware.get_specification().get_semantic_spec() try: @@ -18,7 +18,7 @@ class EasyChecker(LearnwareChecker): except Exception as e: traceback.print_exc() logger.warning(f"The learnware [{learnware.id}] is instantiated failed! Due to {e}") - return cls.NONUSABLE_LEARNWARE + return self.NONUSABLE_LEARNWARE try: learnware_model = learnware.get_model() @@ -35,7 +35,7 @@ class EasyChecker(LearnwareChecker): if stat_spec is not None: if stat_spec.get_z().shape[1:] != input_shape: logger.warning(f"The learnware [{learnware.id}] input dimension mismatch with stat specification") - return cls.NONUSABLE_LEARNWARE + return self.NONUSABLE_LEARNWARE pass inputs = np.random.randn(10, *input_shape) @@ -52,21 +52,21 @@ class EasyChecker(LearnwareChecker): outputs = outputs.detach().cpu().numpy() if not isinstance(outputs, np.ndarray): logger.warning(f"The learnware [{learnware.id}] output must be np.ndarray or torch.Tensor") - return cls.NONUSABLE_LEARNWARE + return self.NONUSABLE_LEARNWARE # check output shape output_dim = int(semantic_spec["Output"]["Dimension"]) if outputs[0].shape[0] != output_dim: logger.warning(f"The learnware [{learnware.id}] input and output dimention is error") - return cls.NONUSABLE_LEARNWARE + return self.NONUSABLE_LEARNWARE pass else: if outputs.shape[1:] != learnware_model.output_shape: logger.warning(f"The learnware [{learnware.id}] input and output dimention is error") - return cls.NONUSABLE_LEARNWARE + return self.NONUSABLE_LEARNWARE except Exception as e: logger.warning(f"The learnware [{learnware.id}] prediction is not avaliable! Due to {repr(e)}") - return cls.NONUSABLE_LEARNWARE + return self.NONUSABLE_LEARNWARE - return cls.USABLE_LEARWARE + return self.USABLE_LEARWARE From d77004063d9032fdbf8971d1e300579605d5f404 Mon Sep 17 00:00:00 2001 From: bxdd Date: Wed, 25 Oct 2023 11:11:40 +0800 Subject: [PATCH 04/35] [MNT] add learnware market interfacve --- learnware/market/base.py | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/learnware/market/base.py b/learnware/market/base.py index 33774d5..737bcb5 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -193,6 +193,11 @@ class BaseMarket: raise NotImplementedError("get semantic spec list is not implemented") + def get_learnware_ids(self) -> List[str]: + raise NotImplementedError("get_learnware_ids is not implemented") + + + class LearnwareOrganizer: def __init__(self, market_id): self.market_id = market_id From 85fc13493272683842a09ec954a894d2db0522ec Mon Sep 17 00:00:00 2001 From: bxdd Date: Thu, 26 Oct 2023 14:18:36 +0800 Subject: [PATCH 05/35] [MNT] add module, implement easy market --- learnware/market/base.py | 16 +-- learnware/market/easy/checker.py | 3 +- learnware/market/easy/database_ops.py | 175 ++++++++++++++++++++++++++ learnware/market/easy/organizer.py | 59 ++++++++- learnware/market/easy/searcher.py | 69 +++++++++- learnware/market/module.py | 0 6 files changed, 310 insertions(+), 12 deletions(-) create mode 100644 learnware/market/easy/database_ops.py create mode 100644 learnware/market/module.py diff --git a/learnware/market/base.py b/learnware/market/base.py index 737bcb5..5f3dc56 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -61,16 +61,16 @@ class BaseMarket: self.learnware_searcher = LearnwareSearcher() if searcher is None else searcher self.learnware_searcher.reset(organizer=self.learnware_organizer) - def reload_market(self, *args, **kwargs) -> bool: + def reload_market(self, **kwargs) -> bool: """Reload the market when server restared. Returns ------- bool A flag indicating whether the market is reload successfully. """ - self.learnware_organizer.reload_market(*args, **kwargs) + self.learnware_organizer.reload_market(**kwargs) - def check_learnware(self, learnware: Learnware) -> bool: + def check_learnware(self, learnware: Learnware, **kwargs) -> bool: """Check the utility of a learnware Parameters @@ -82,7 +82,7 @@ class BaseMarket: bool A flag indicating whether the learnware can be accepted. """ - return self.learnware_checker(learnware) + return self.learnware_checker(learnware, **kwargs) def add_learnware(self, zip_path: str, semantic_spec: dict, **kwargs) -> Tuple[str, bool]: """Add a learnware into the market. @@ -111,7 +111,7 @@ class BaseMarket: """ return self.learnware_organizer.add_learnware(zip_path, semantic_spec, **kwargs) - def search_learnware(self, user_info: BaseUserInfo) -> Tuple[Any, List[Learnware]]: + def search_learnware(self, user_info: BaseUserInfo, **kwargs) -> Tuple[Any, List[Learnware]]: """Search Learnware based on user_info Parameters @@ -128,7 +128,7 @@ class BaseMarket: - second is a list of matched learnwares """ - return self.learnware_searcher(user_info=user_info) + return self.learnware_searcher(user_info, *args, **kwargs) def get_learnware_by_ids(self, id: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """ @@ -149,7 +149,7 @@ class BaseMarket: """ return self.learnware_organizer.get_learnware_by_ids(id) - def delete_learnware(self, id: str) -> bool: + def delete_learnware(self, id: str, *args, **kwargs) -> bool: """Delete a learnware from market Parameters @@ -167,7 +167,7 @@ class BaseMarket: Exception Raise an excpetion when given id is NOT found in learnware list """ - return self.learnware_organizer.delete_learnware(id) + return self.learnware_organizer.delete_learnware(id, **kwargs) def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, **kwargs) -> bool: """ diff --git a/learnware/market/easy/checker.py b/learnware/market/easy/checker.py index 50fac94..efb48b2 100644 --- a/learnware/market/easy/checker.py +++ b/learnware/market/easy/checker.py @@ -1,13 +1,12 @@ import traceback -from ..base import LearnwareChecker, LearnwareOrganizer +from ..base import LearnwareChecker from ...logger import get_module_logger logger = get_module_logger("easy_checker", "INFO") class EasyChecker(LearnwareChecker): - def __call__(self, learnware): semantic_spec = learnware.get_specification().get_semantic_spec() diff --git a/learnware/market/easy/database_ops.py b/learnware/market/easy/database_ops.py new file mode 100644 index 0000000..48ed0fb --- /dev/null +++ b/learnware/market/easy/database_ops.py @@ -0,0 +1,175 @@ +from sqlalchemy.ext.declarative import declarative_base +from sqlalchemy import create_engine, text +from sqlalchemy import Column, Text, String +import os +import json +from ...learnware import get_learnware_from_dirpath +from ...logger import get_module_logger + +logger = get_module_logger("database") +DeclarativeBase = declarative_base() + + +class Learnware(DeclarativeBase): + __tablename__ = "tb_learnware" + + id = Column(String(10), primary_key=True, nullable=False) + semantic_spec = Column(Text, nullable=False) + zip_path = Column(Text, nullable=False) + folder_path = Column(Text, nullable=False) + use_flag = Column(Text, nullable=False) + + pass + + +class DatabaseOperations(object): + def __init__(self, url: str, database_name: str): + if url.startswith("sqlite"): + url = os.path.join(url, f"{database_name}.db") + else: + url = f"{url}/{database_name}" + pass + + self.url = url + self.create_database_if_not_exists(url) + + pass + + def create_database_if_not_exists(self, url): + database_exists = True + + if url.startswith("sqlite"): + # it is sqlite + start = url.find(":///") + path = url[start + 4 :] + if os.path.exists(path): + database_exists = True + pass + else: + database_exists = False + os.makedirs(os.path.dirname(path), exist_ok=True) + pass + pass + elif self.url.startswith("postgresql"): + # it is postgresql + dbname_start = url.rfind("/") + dbname = url[dbname_start + 1 :] + url_no_dbname = url[:dbname_start] + "/postgres" + engine = create_engine(url_no_dbname) + + with engine.connect() as conn: + result = conn.execute(text("SELECT datname FROM pg_database;")) + db_list = set() + + for row in result.fetchall(): + db_list.add(row[0].lower()) + pass + + if dbname.lower() not in db_list: + database_exists = False + conn.execution_options(isolation_level="AUTOCOMMIT").execute( + text("CREATE DATABASE {0};".format(dbname)) + ) + pass + else: + database_exists = True + pass + pass + engine.dispose() + pass + else: + raise Exception(f"Unsupported database url: {self.url}") + pass + + self.engine = create_engine(url, future=True) + + if not database_exists: + DeclarativeBase.metadata.create_all(self.engine) + pass + pass + + def clear_learnware_table(self): + with self.engine.connect() as conn: + conn.execute(text("DELETE FROM tb_learnware;")) + conn.commit() + pass + pass + + def add_learnware(self, id: str, semantic_spec: dict, zip_path, folder_path, use_flag: str): + with self.engine.connect() as conn: + semantic_spec_str = json.dumps(semantic_spec) + conn.execute( + text( + ( + "INSERT INTO tb_learnware (id, semantic_spec, zip_path, folder_path, use_flag)" + "VALUES (:id, :semantic_spec, :zip_path, :folder_path, :use_flag);" + ) + ), + dict( + id=id, + semantic_spec=semantic_spec_str, + zip_path=zip_path, + folder_path=folder_path, + use_flag=use_flag, + ), + ) + conn.commit() + pass + pass + + def update_learnware_semantic_spec(self, learnware_id: str, semantic_spec: dict): + with self.engine.connect() as conn: + semantic_spec_str = json.dumps(semantic_spec) + conn.execute( + text("UPDATE tb_learnware SET semantic_spec=:semantic_spec WHERE id=:id;"), + dict(id=learnware_id, semantic_spec=semantic_spec_str), + ) + conn.commit() + pass + pass + + def delete_learnware(self, id: str): + with self.engine.connect() as conn: + conn.execute(text("DELETE FROM tb_learnware WHERE id=:id;"), dict(id=id)) + conn.commit() + pass + pass + + def update_learnware_semantic_specification(self, id: str, semantic_spec: dict): + with self.engine.connect() as conn: + semantic_spec_str = json.dumps(semantic_spec) + r = conn.execute( + text("UPDATE tb_learnware SET semantic_spec=:semantic_spec WHERE id=:id;"), + dict(id=id, semantic_spec=semantic_spec_str), + ) + conn.commit() + pass + pass + + def load_market(self): + with self.engine.connect() as conn: + cursor = conn.execute(text("SELECT id, semantic_spec, zip_path, folder_path, use_flag FROM tb_learnware;")) + + learnware_list = {} + zip_list = {} + folder_list = {} + max_count = 0 + + for id, semantic_spec, zip_path, folder_path, use_flag in cursor: + id = id.strip() + semantic_spec_dict = json.loads(semantic_spec) + new_learnware = get_learnware_from_dirpath( + id=id, semantic_spec=semantic_spec_dict, learnware_dirpath=folder_path + ) + logger.info(f"Load learnware: {id}") + learnware_list[id] = new_learnware + # assert new_learnware is not None + zip_list[id] = zip_path + folder_list[id] = folder_path + max_count = max(max_count, int(id)) + pass + + return learnware_list, zip_list, folder_list, max_count + 1 + pass + + pass diff --git a/learnware/market/easy/organizer.py b/learnware/market/easy/organizer.py index cdbaf69..02e1db7 100644 --- a/learnware/market/easy/organizer.py +++ b/learnware/market/easy/organizer.py @@ -1,4 +1,24 @@ +import os +import json +import copy +import torch +import zipfile import traceback +import numpy as np +import pandas as pd +from rapidfuzz import fuzz +from cvxopt import solvers, matrix +from shutil import copyfile, rmtree +from typing import Tuple, Any, List, Union, Dict + +from ..base import BaseMarket, BaseUserInfo +from ..database_ops import DatabaseOperations + +from ... import utils +from ...config import C as conf +from ...logger import get_module_logger +from ...learnware import Learnware, get_learnware_from_dirpath +from ...specification import RKMEStatSpecification, Specification from ..base import LearnwareOrganizer from ...logger import get_module_logger @@ -8,4 +28,41 @@ logger = get_module_logger("easy_organizer") class EasyOrganizer(LearnwareOrganizer): - \ No newline at end of file + def reset(self, market_id): + self.market_id = market_id + self.reload_market() + + def reload_market(self, rebuild=False) -> bool: + """Reload the learnware organizer when server restared. + + Returns + ------- + bool + A flag indicating whether the market is reload successfully. + """ + + self.market_store_path = os.path.join(conf.market_root_path, self.market_id) + self.learnware_pool_path = os.path.join(self.market_store_path, "learnware_pool") + self.learnware_zip_pool_path = os.path.join(self.learnware_pool_path, "zips") + self.learnware_folder_pool_path = os.path.join(self.learnware_pool_path, "unzipped_learnwares") + self.learnware_list = {} # id: Learnware + self.learnware_zip_list = {} + self.learnware_folder_list = {} + self.count = 0 + self.semantic_spec_list = conf.semantic_specs + self.dbops = DatabaseOperations(conf.database_url, "market_" + self.market_id) + + if rebuild: + logger.warning("Warning! You are trying to clear current database!") + try: + self.dbops.clear_learnware_table() + rmtree(self.learnware_pool_path) + except: + pass + + os.makedirs(self.learnware_pool_path, exist_ok=True) + os.makedirs(self.learnware_zip_pool_path, exist_ok=True) + os.makedirs(self.learnware_folder_pool_path, exist_ok=True) + self.learnware_list, self.learnware_zip_list, self.learnware_folder_list, self.count = self.dbops.load_market() + + \ No newline at end of file diff --git a/learnware/market/easy/searcher.py b/learnware/market/easy/searcher.py index eb97a7f..427f51c 100644 --- a/learnware/market/easy/searcher.py +++ b/learnware/market/easy/searcher.py @@ -1,8 +1,75 @@ +from typing import Tuple, List + from ..base import LearnwareSearcher from ...logger import get_module_logger +from ...learnware import Learnware +from ...market import BaseUserInfo logger = get_module_logger('easy_seacher') class EasySearcher(LearnwareSearcher): - pass + + def __call__(self, user_info: BaseUserInfo, max_search_num: int = 5, search_method: str = "greedy") -> Tuple[List[float], List[Learnware], float, List[Learnware]]: + """Search learnwares based on user_info + + Parameters + ---------- + user_info : BaseUserInfo + user_info contains semantic_spec and stat_info + max_search_num : int + The maximum number of the returned learnwares + + Returns + ------- + Tuple[List[float], List[Learnware], float, List[Learnware]] + the first is the sorted list of rkme dist + the second is the sorted list of Learnware (single) by the rkme dist + the third is the score of Learnware (mixture) + the fourth is the list of Learnware (mixture), the size is search_num + """ + learnware_list = [self.learnware_list[key] for key in self.learnware_list] + # learnware_list = self._search_by_semantic_spec_exact(learnware_list, user_info) + # if len(learnware_list) == 0: + learnware_list = self._search_by_semantic_spec_fuzz(learnware_list, user_info) + + if "RKMEStatSpecification" not in user_info.stat_info: + return None, learnware_list, 0.0, None + elif len(learnware_list) == 0: + return [], [], 0.0, [] + else: + user_rkme = user_info.stat_info["RKMEStatSpecification"] + learnware_list = self._filter_by_rkme_spec_dimension(learnware_list, user_rkme) + logger.info(f"After filter by rkme dimension, learnware_list length is {len(learnware_list)}") + + sorted_dist_list, single_learnware_list = self._search_by_rkme_spec_single(learnware_list, user_rkme) + if search_method == "auto": + mixture_dist, weight_list, mixture_learnware_list = self._search_by_rkme_spec_mixture_auto( + learnware_list, user_rkme, max_search_num + ) + elif search_method == "greedy": + mixture_dist, weight_list, mixture_learnware_list = self._search_by_rkme_spec_mixture_greedy( + learnware_list, user_rkme, max_search_num + ) + else: + logger.warning("f{search_method} not supported!") + mixture_dist = None + weight_list = [] + mixture_learnware_list = [] + + if mixture_dist is None: + sorted_score_list = self._convert_dist_to_score(sorted_dist_list) + mixture_score = None + else: + merge_score_list = self._convert_dist_to_score(sorted_dist_list + [mixture_dist]) + sorted_score_list = merge_score_list[:-1] + mixture_score = merge_score_list[-1] + + logger.info(f"After search by rkme spec, learnware_list length is {len(learnware_list)}") + # filter learnware with low score + sorted_score_list, single_learnware_list = self._filter_by_rkme_spec_single( + sorted_score_list, single_learnware_list + ) + + logger.info(f"After filter by rkme spec, learnware_list length is {len(learnware_list)}") + return sorted_score_list, single_learnware_list, mixture_score, mixture_learnware_list \ No newline at end of file diff --git a/learnware/market/module.py b/learnware/market/module.py new file mode 100644 index 0000000..e69de29 From 01c96c1abea0f1cff65a23cd23caa94f3cb6514f Mon Sep 17 00:00:00 2001 From: bxdd Date: Thu, 26 Oct 2023 15:56:42 +0800 Subject: [PATCH 06/35] [MNT] rename basemarket to learnware market --- docs/references/api.rst | 2 +- learnware/market/__init__.py | 2 +- learnware/market/anchor.py | 8 +- learnware/market/base.py | 61 ++++---- learnware/market/easy.py | 6 +- learnware/market/easy/organizer.py | 226 ++++++++++++++++++++++++++++- learnware/market/evolve.py | 6 +- 7 files changed, 265 insertions(+), 46 deletions(-) diff --git a/docs/references/api.rst b/docs/references/api.rst index ebd276f..a2f723b 100644 --- a/docs/references/api.rst +++ b/docs/references/api.rst @@ -11,7 +11,7 @@ Here you can find all ``learnware`` interfaces. Market ==================== -.. autoclass:: learnware.market.BaseMarket +.. autoclass:: learnware.market.LearnwareMarket :members: .. autoclass:: learnware.market.EasyMarket diff --git a/learnware/market/__init__.py b/learnware/market/__init__.py index 1620a5e..bf18990 100644 --- a/learnware/market/__init__.py +++ b/learnware/market/__init__.py @@ -1,5 +1,5 @@ from .anchor import AnchoredUserInfo, AnchoredMarket -from .base import BaseUserInfo, BaseMarket +from .base import BaseUserInfo, LearnwareMarket from .evolve_anchor import EvolvedAnchoredMarket from .evolve import EvolvedMarket from .easy import EasyMarket diff --git a/learnware/market/anchor.py b/learnware/market/anchor.py index 79d5443..bd912f3 100644 --- a/learnware/market/anchor.py +++ b/learnware/market/anchor.py @@ -2,7 +2,7 @@ import os from typing import Tuple, Any, List, Union, Dict from ..learnware import Learnware -from .base import BaseMarket, BaseUserInfo +from .base import LearnwareMarket, BaseUserInfo class AnchoredUserInfo(BaseUserInfo): @@ -42,12 +42,12 @@ class AnchoredUserInfo(BaseUserInfo): self.stat_info[name] = item -class AnchoredMarket(BaseMarket): - """Add the anchor design to the BaseMarket +class AnchoredMarket(LearnwareMarket): + """Add the anchor design to the LearnwareMarket Parameters ---------- - BaseMarket : _type_ + LearnwareMarket : _type_ Basic market version """ diff --git a/learnware/market/base.py b/learnware/market/base.py index 5f3dc56..12c8f9f 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -43,7 +43,7 @@ class BaseUserInfo: return self.stat_info.get(name, None) -class BaseMarket: +class LearnwareMarket: """Base interface for market, it provide the interface of search/add/detele/update learnwares""" def __init__( @@ -55,9 +55,9 @@ class BaseMarket: ): self.market_id = market_id self.learnware_organizer = LearnwareOrganizer() if organizer is None else organizer - self.learnware_organizer.reset(market_id=market_id) self.learnware_checker = LearnwareChecker() if checker is None else checker self.learnware_checker.reset(organizer=self.learnware_organizer) + self.learnware_organizer.reset(market_id=market_id, checker=self.learnware_checker) self.learnware_searcher = LearnwareSearcher() if searcher is None else searcher self.learnware_searcher.reset(organizer=self.learnware_organizer) @@ -128,26 +128,7 @@ class BaseMarket: - second is a list of matched learnwares """ - return self.learnware_searcher(user_info, *args, **kwargs) - - def get_learnware_by_ids(self, id: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: - """ - Get Learnware from market by id - - Parameters - ---------- - id : Union[str, List[str]] - Given one id or a list of ids as target. - - Returns - ------- - Union[Learnware, List[Learnware]] - Return a Learnware object or a list of Learnware objects based on the type of input param. - - - The returned items are search results. - - 'None' indicating the target id not found. - """ - return self.learnware_organizer.get_learnware_by_ids(id) + return self.learnware_searcher(user_info, **kwargs) def delete_learnware(self, id: str, *args, **kwargs) -> bool: """Delete a learnware from market @@ -191,19 +172,35 @@ class BaseMarket: """ raise NotImplementedError("get semantic spec list is not implemented") + + def get_learnware_by_ids(self, id: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: + """ + Get Learnware from market by id + Parameters + ---------- + id : Union[str, List[str]] + Given one id or a list of ids as target. - def get_learnware_ids(self) -> List[str]: - raise NotImplementedError("get_learnware_ids is not implemented") - - + Returns + ------- + Union[Learnware, List[Learnware]] + Return a Learnware object or a list of Learnware objects based on the type of input param. + + - The returned items are search results. + - 'None' indicating the target id not found. + """ + return self.learnware_organizer.get_learnware_by_ids(id) + + def class LearnwareOrganizer: - def __init__(self, market_id): - self.market_id = market_id + def __init__(self, market_id, organizer: 'LearnwareOrganizer' = None): + self.reset(market_id=market_id, organizer=organizer) - def reset(self, market_id): + def reset(self, market_id, organizer: 'LearnwareOrganizer' ): self.market_id = market_id + self.organizer = organizer def reload_market(self) -> bool: """Reload the learnware organizer when server restared. @@ -327,6 +324,12 @@ class LearnwareOrganizer: """ raise NotImplementedError("get_learnware_path_by_ids is not implemented") + def get_learnware_ids(self, top:int = None): + if top is None: + return list(self.learnware_list.keys()) + else: + return list(self.learnware_list.keys())[:top] + class LearnwareSearcher: def __init__(self, organizer: LearnwareOrganizer = None): self.learnware_oganizer = organizer diff --git a/learnware/market/easy.py b/learnware/market/easy.py index 26f8fb0..c577a82 100644 --- a/learnware/market/easy.py +++ b/learnware/market/easy.py @@ -11,7 +11,7 @@ from cvxopt import solvers, matrix from shutil import copyfile, rmtree from typing import Tuple, Any, List, Union, Dict -from .base import BaseMarket, BaseUserInfo +from .base import LearnwareMarket, BaseUserInfo from .database_ops import DatabaseOperations from .. import utils @@ -24,8 +24,8 @@ from ..specification import RKMEStatSpecification, Specification logger = get_module_logger("market", "INFO") -class EasyMarket(BaseMarket): - """EasyMarket provide an easy and simple implementation for BaseMarket +class EasyMarket(LearnwareMarket): + """EasyMarket provide an easy and simple implementation for LearnwareMarket - EasyMarket stores learnwares with file system and database - EasyMarket search the learnwares with the match of semantical tag and the statistical RKME - EasyMarket does not support the search between heterogeneous features learnwars diff --git a/learnware/market/easy/organizer.py b/learnware/market/easy/organizer.py index 02e1db7..d3a70f8 100644 --- a/learnware/market/easy/organizer.py +++ b/learnware/market/easy/organizer.py @@ -11,7 +11,7 @@ from cvxopt import solvers, matrix from shutil import copyfile, rmtree from typing import Tuple, Any, List, Union, Dict -from ..base import BaseMarket, BaseUserInfo +from ..base import LearnwareMarket, BaseUserInfo from ..database_ops import DatabaseOperations from ... import utils @@ -20,7 +20,7 @@ from ...logger import get_module_logger from ...learnware import Learnware, get_learnware_from_dirpath from ...specification import RKMEStatSpecification, Specification -from ..base import LearnwareOrganizer +from ..base import LearnwareOrganizer, LearnwareChecker from ...logger import get_module_logger logger = get_module_logger("easy_organizer") @@ -28,9 +28,9 @@ logger = get_module_logger("easy_organizer") class EasyOrganizer(LearnwareOrganizer): - def reset(self, market_id): + def reset(self, market_id, rebuild=False): self.market_id = market_id - self.reload_market() + self.reload_market(rebuild=rebuild) def reload_market(self, rebuild=False) -> bool: """Reload the learnware organizer when server restared. @@ -64,5 +64,221 @@ class EasyOrganizer(LearnwareOrganizer): os.makedirs(self.learnware_zip_pool_path, exist_ok=True) os.makedirs(self.learnware_folder_pool_path, exist_ok=True) self.learnware_list, self.learnware_zip_list, self.learnware_folder_list, self.count = self.dbops.load_market() + + + def add_learnware(self, zip_path: str, semantic_spec: dict, learnware_id: str = None, check: bool = False) -> Tuple[str, bool]: + """Add a learnware into the market. + + .. note:: + + Given a prediction of a certain time, all signals before this time will be prepared well. + + + Parameters + ---------- + zip_path : str + Filepath for learnware model, a zipped file. + semantic_spec : dict + semantic_spec for new learnware, in dictionary format. + + Returns + ------- + Tuple[str, int] + - str indicating model_id + - int indicating what the flag of learnware is added. + + """ + semantic_spec = copy.deepcopy(semantic_spec) + + if not os.path.exists(zip_path): + logger.warning("Zip Path NOT Found! Fail to add learnware.") + return None, self.INVALID_LEARNWARE + + try: + if len(semantic_spec["Data"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please choose Data.") + return None, self.INVALID_LEARNWARE + if len(semantic_spec["Task"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please choose Task.") + return None, self.INVALID_LEARNWARE + if len(semantic_spec["Library"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please choose Device.") + return None, self.INVALID_LEARNWARE + if len(semantic_spec["Name"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please provide Name.") + return None, self.INVALID_LEARNWARE + if len(semantic_spec["Description"]["Values"]) == 0 and len(semantic_spec["Scenario"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please provide Scenario or Description.") + return None, self.INVALID_LEARNWARE + if ( + semantic_spec["Data"]["Type"] != "Class" + or semantic_spec["Task"]["Type"] != "Class" + or semantic_spec["Library"]["Type"] != "Class" + or semantic_spec["Scenario"]["Type"] != "Tag" + or semantic_spec["Name"]["Type"] != "String" + or semantic_spec["Description"]["Type"] != "String" + ): + logger.warning("Illegal semantic specification, please provide the right type.") + return None, self.INVALID_LEARNWARE + except: + print(semantic_spec) + logger.warning("Illegal semantic specification, some keys are missing.") + return None, self.INVALID_LEARNWARE + + logger.info("Get new learnware from %s" % (zip_path)) + if learnware_id is not None: + id = learnware_id + else: + id = "%08d" % (self.count) + target_zip_dir = os.path.join(self.learnware_zip_pool_path, "%s.zip" % (id)) + target_folder_dir = os.path.join(self.learnware_folder_pool_path, id) + copyfile(zip_path, target_zip_dir) + + with zipfile.ZipFile(target_zip_dir, "r") as z_file: + z_file.extractall(target_folder_dir) + logger.info("Learnware move to %s, and unzip to %s" % (target_zip_dir, target_folder_dir)) + + try: + new_learnware = get_learnware_from_dirpath( + id=id, semantic_spec=semantic_spec, learnware_dirpath=target_folder_dir + ) + except: + try: + os.remove(target_zip_dir) + rmtree(target_folder_dir) + except: + pass + return None, self.INVALID_LEARNWARE + + if new_learnware is None: + return None, self.INVALID_LEARNWARE + + if check and self.checker + + self.dbops.add_learnware( + id=id, + semantic_spec=semantic_spec, + zip_path=target_zip_dir, + folder_path=target_folder_dir, + use_flag=LearnwareChecker.USABLE_LEARWARE, + ) + + self.learnware_list[id] = new_learnware + self.learnware_zip_list[id] = target_zip_dir + self.learnware_folder_list[id] = target_folder_dir + self.count += 1 + return id, LearnwareChecker.USABLE_LEARWARE + + def add_learnware(self, zip_path: str, semantic_spec: dict) -> Tuple[str, bool]: + """Add a learnware into the market. + + .. note:: + + Given a prediction of a certain time, all signals before this time will be prepared well. + + + Parameters + ---------- + zip_path : str + Filepath for learnware model, a zipped file. + semantic_spec : dict + semantic_spec for new learnware, in dictionary format. + + Returns + ------- + Tuple[str, int] + - str indicating model_id + - int indicating what the flag of learnware is added. + + """ + semantic_spec = copy.deepcopy(semantic_spec) + + if not os.path.exists(zip_path): + logger.warning("Zip Path NOT Found! Fail to add learnware.") + return None, self.INVALID_LEARNWARE + + try: + if len(semantic_spec["Data"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please choose Data.") + return None, self.INVALID_LEARNWARE + if len(semantic_spec["Task"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please choose Task.") + return None, self.INVALID_LEARNWARE + if len(semantic_spec["Library"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please choose Device.") + return None, self.INVALID_LEARNWARE + if len(semantic_spec["Name"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please provide Name.") + return None, self.INVALID_LEARNWARE + if len(semantic_spec["Description"]["Values"]) == 0 and len(semantic_spec["Scenario"]["Values"]) == 0: + logger.warning("Illegal semantic specification, please provide Scenario or Description.") + return None, self.INVALID_LEARNWARE + if ( + semantic_spec["Data"]["Type"] != "Class" + or semantic_spec["Task"]["Type"] != "Class" + or semantic_spec["Library"]["Type"] != "Class" + or semantic_spec["Scenario"]["Type"] != "Tag" + or semantic_spec["Name"]["Type"] != "String" + or semantic_spec["Description"]["Type"] != "String" + ): + logger.warning("Illegal semantic specification, please provide the right type.") + return None, self.INVALID_LEARNWARE + except: + logger.info(f"Semantic specification: {semantic_spec}") + logger.warning("Illegal semantic specification, some keys are missing.") + return None, self.INVALID_LEARNWARE + + logger.info("Get new learnware from %s" % (zip_path)) + id = "%08d" % (self.count) + target_zip_dir = os.path.join(self.learnware_zip_pool_path, "%s.zip" % (id)) + target_folder_dir = os.path.join(self.learnware_folder_pool_path, id) + copyfile(zip_path, target_zip_dir) + + with zipfile.ZipFile(target_zip_dir, "r") as z_file: + z_file.extractall(target_folder_dir) + logger.info("Learnware move to %s, and unzip to %s" % (target_zip_dir, target_folder_dir)) + + try: + new_learnware = get_learnware_from_dirpath( + id=id, semantic_spec=semantic_spec, learnware_dirpath=target_folder_dir + ) + except: + try: + os.remove(target_zip_dir) + rmtree(target_folder_dir) + except: + pass + return None, self.INVALID_LEARNWARE + + if new_learnware is None: + return None, self.INVALID_LEARNWARE + + check_flag = self.check_learnware(new_learnware) + + self.dbops.add_learnware( + id=id, + semantic_spec=semantic_spec, + zip_path=target_zip_dir, + folder_path=target_folder_dir, + use_flag=check_flag, + ) + + self.learnware_list[id] = new_learnware + self.learnware_zip_list[id] = target_zip_dir + self.learnware_folder_list[id] = target_folder_dir + self.count += 1 + return id, check_flag + + + def get_learnware_ids(self, top:int = None): + if top is None: + return list(self.learnware_list.keys()) + else: + return list(self.learnware_list.keys())[:top] - \ No newline at end of file + + def get_learnwares(self, top:int = None): + if top is None: + return list(self.learnware_list.values()) + else: + return list(self.learnware_list.values())[:top] \ No newline at end of file diff --git a/learnware/market/evolve.py b/learnware/market/evolve.py index 8912700..e9e5cc3 100644 --- a/learnware/market/evolve.py +++ b/learnware/market/evolve.py @@ -1,16 +1,16 @@ from typing import Tuple, Any, List, Union, Dict -from .base import BaseMarket +from .base import LearnwareMarket from ..learnware import Learnware from ..specification import BaseStatSpecification -class EvolvedMarket(BaseMarket): +class EvolvedMarket(LearnwareMarket): """Organize learnwares and enable them to continuously evolve Parameters ---------- - BaseMarket : _type_ + LearnwareMarket : _type_ Basic market version """ From 50551d42656ebbc2deb2d1f80de1bdc06695cd71 Mon Sep 17 00:00:00 2001 From: bxdd Date: Thu, 26 Oct 2023 22:32:39 +0800 Subject: [PATCH 07/35] [MNT] maintain use_flags in easy market --- learnware/market/base.py | 192 ++++++------------ learnware/market/database_ops.py | 18 +- learnware/market/easy.py | 4 +- learnware/market/easy/checker.py | 2 + learnware/market/easy/database_ops.py | 31 +-- learnware/market/easy/organizer.py | 275 +++++++++++++++----------- 6 files changed, 241 insertions(+), 281 deletions(-) diff --git a/learnware/market/base.py b/learnware/market/base.py index 12c8f9f..2bec206 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -43,6 +43,10 @@ class BaseUserInfo: return self.stat_info.get(name, None) +class BaseSearchResult: + + pass + class LearnwareMarket: """Base interface for market, it provide the interface of search/add/detele/update learnwares""" @@ -62,145 +66,43 @@ class LearnwareMarket: self.learnware_searcher.reset(organizer=self.learnware_organizer) def reload_market(self, **kwargs) -> bool: - """Reload the market when server restared. - Returns - ------- - bool - A flag indicating whether the market is reload successfully. - """ self.learnware_organizer.reload_market(**kwargs) def check_learnware(self, learnware: Learnware, **kwargs) -> bool: - """Check the utility of a learnware - - Parameters - ---------- - learnware : Learnware - - Returns - ------- - bool - A flag indicating whether the learnware can be accepted. - """ return self.learnware_checker(learnware, **kwargs) def add_learnware(self, zip_path: str, semantic_spec: dict, **kwargs) -> Tuple[str, bool]: - """Add a learnware into the market. - - .. note:: - - Given a prediction of a certain time, all signals before this time will be prepared well. - - Parameters - ---------- - zip_path : str - Filepath for learnware model, a zipped file. - semantic_spec : dict - semantic_spec for new learnware, in dictionary format. - - Returns - ------- - Tuple[str, bool] - str indicating model_id, bool indicating whether the learnware is added successfully. - - Raises - ------ - FileNotFoundError - file for model or statistical specification not found - - """ return self.learnware_organizer.add_learnware(zip_path, semantic_spec, **kwargs) def search_learnware(self, user_info: BaseUserInfo, **kwargs) -> Tuple[Any, List[Learnware]]: - """Search Learnware based on user_info - - Parameters - ---------- - user_info : BaseUserInfo - user_info with emantic specifications and statistical information - - Returns - ------- - Tuple[Any, List[Any]] - return two items: - - - first is recommended combination, None when no recommended combination is calculated or statistical specification is not provided. - - second is a list of matched learnwares - """ - return self.learnware_searcher(user_info, **kwargs) - def delete_learnware(self, id: str, *args, **kwargs) -> bool: - """Delete a learnware from market - - Parameters - ---------- - id : str - id of learnware to be deleted - - Returns - ------- - bool - True if the target learnware is deleted successfully. - - Raises - ------ - Exception - Raise an excpetion when given id is NOT found in learnware list - """ + def delete_learnware(self, id: str, **kwargs) -> bool: return self.learnware_organizer.delete_learnware(id, **kwargs) def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, **kwargs) -> bool: - """ - Update Learnware with id and content to be updated. - Empty interface. TODO - - Parameters - ---------- - id : str - id of target learnware. - """ return self.learnware_organizer.update_learnware(id, zip_path=zip_path, semantic_spec=semantic_spec, **kwargs) - def get_semantic_spec_list(self) -> dict: - """Return all semantic specifications available - - Returns - ------- - dict - All emantic specifications in dictionary format - - """ - raise NotImplementedError("get semantic spec list is not implemented") + def get_learnware_ids(self, top:int = None, **kwargs): + return self.learnware_organizer.get_learnware_ids(top, **kwargs) + - def get_learnware_by_ids(self, id: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: - """ - Get Learnware from market by id - - Parameters - ---------- - id : Union[str, List[str]] - Given one id or a list of ids as target. - - Returns - ------- - Union[Learnware, List[Learnware]] - Return a Learnware object or a list of Learnware objects based on the type of input param. - - - The returned items are search results. - - 'None' indicating the target id not found. - """ - return self.learnware_organizer.get_learnware_by_ids(id) + def get_learnwares(self, top:int = None, **kwargs): + return self.learnware_organizer.get_learnwares(top, **kwargs) + + def get_learnware_path_by_ids(self, ids: Union[str, List[str]], **kwargs) -> Union[Learnware, List[Learnware]]: + raise self.learnware_organizer.get_learnware_path_by_ids(ids, **kwargs) - def + def get_learnware_by_ids(self, id: Union[str, List[str]], **kwargs) -> Union[Learnware, List[Learnware]]: + return self.learnware_organizer.get_learnware_by_ids(id, **kwargs) class LearnwareOrganizer: - def __init__(self, market_id, organizer: 'LearnwareOrganizer' = None): - self.reset(market_id=market_id, organizer=organizer) + def __init__(self, market_id, checker: 'LearnwareChecker' = None): + self.reset(market_id=market_id, checker=checker) - def reset(self, market_id, organizer: 'LearnwareOrganizer' ): + def reset(self, market_id, checker: 'LearnwareChecker', **kwargs): self.market_id = market_id - self.organizer = organizer + self.organizer = checker def reload_market(self) -> bool: """Reload the learnware organizer when server restared. @@ -267,7 +169,6 @@ class LearnwareOrganizer: def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, **kwargs) -> bool: """ Update Learnware with id and content to be updated. - Empty interface. TODO Parameters ---------- @@ -276,17 +177,6 @@ class LearnwareOrganizer: """ raise NotImplementedError("update learnware is Not Implemented") - def get_semantic_spec_list(self) -> dict: - """Return all semantic specifications available - - Returns - ------- - dict - All emantic specifications in dictionary format - - """ - raise NotImplementedError("get semantic spec list is not implemented") - def get_learnware_by_ids(self, id: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """ Get Learnware from market by id @@ -324,11 +214,36 @@ class LearnwareOrganizer: """ raise NotImplementedError("get_learnware_path_by_ids is not implemented") - def get_learnware_ids(self, top:int = None): - if top is None: - return list(self.learnware_list.keys()) - else: - return list(self.learnware_list.keys())[:top] + def get_learnware_ids(self, top:int = None) -> List[str]: + """get the list of learnware ids + + Parameters + ---------- + top : int, optional + the first top element to return, by default None + + Raises + ------ + List[str] + the first top ids + """ + raise NotImplementedError("get_learnware_ids is not implemented") + + + def get_learnwares(self, top:int = None) -> List[Learnware]: + """get the list of learnwares + + Parameters + ---------- + top : int, optional + the first top element to return, by default None + + Raises + ------ + List[Learnware] + the first top learnwares + """ + raise NotImplementedError("get_learnwares is not implemented") class LearnwareSearcher: def __init__(self, organizer: LearnwareOrganizer = None): @@ -337,7 +252,14 @@ class LearnwareSearcher: def reset(self, organizer): self.learnware_oganizer = organizer - def __call__(self, user_info: BaseUserInfo): + def __call__(self, user_info: BaseUserInfo) + """Search learnwares based on user_info + + Parameters + ---------- + user_info : BaseUserInfo + user_info contains semantic_spec and stat_info + """ raise NotImplementedError("'__call__' method is not implemented in LearnwareSearcher") diff --git a/learnware/market/database_ops.py b/learnware/market/database_ops.py index f656425..c44aa01 100644 --- a/learnware/market/database_ops.py +++ b/learnware/market/database_ops.py @@ -117,25 +117,25 @@ class DatabaseOperations(object): pass pass - def update_learnware_semantic_spec(self, learnware_id: str, semantic_spec: dict): + def delete_learnware(self, id: str): with self.engine.connect() as conn: - semantic_spec_str = json.dumps(semantic_spec) - conn.execute( - text("UPDATE tb_learnware SET semantic_spec=:semantic_spec WHERE id=:id;"), - dict(id=learnware_id, semantic_spec=semantic_spec_str), - ) + conn.execute(text("DELETE FROM tb_learnware WHERE id=:id;"), dict(id=id)) conn.commit() pass pass - def delete_learnware(self, id: str): + def update_learnware_semantic_specification(self, id: str, semantic_spec: dict): with self.engine.connect() as conn: - conn.execute(text("DELETE FROM tb_learnware WHERE id=:id;"), dict(id=id)) + semantic_spec_str = json.dumps(semantic_spec) + r = conn.execute( + text("UPDATE tb_learnware SET semantic_spec=:semantic_spec WHERE id=:id;"), + dict(id=id, semantic_spec=semantic_spec_str), + ) conn.commit() pass pass - def update_learnware_semantic_specification(self, id: str, semantic_spec: dict): + def update_learnware_use_flag(self, id: str, semantic_spec: dict): with self.engine.connect() as conn: semantic_spec_str = json.dumps(semantic_spec) r = conn.execute( diff --git a/learnware/market/easy.py b/learnware/market/easy.py index c577a82..63a357e 100644 --- a/learnware/market/easy.py +++ b/learnware/market/easy.py @@ -949,11 +949,11 @@ class EasyMarket(LearnwareMarket): logger.warning("Learnware ID '%s' NOT Found!" % (ids)) return None - def update_learnware_semantic_spec(self, learnware_id: str, semantic_spec: dict) -> bool: + def update_learnware_semantic_specification(self, learnware_id: str, semantic_spec: dict) -> bool: """Update Learnware semantic_spec""" # update database - self.dbops.update_learnware_semantic_spec(learnware_id=learnware_id, semantic_spec=semantic_spec) + self.dbops.update_learnware_semantic_specification(learnware_id=learnware_id, semantic_spec=semantic_spec) # update file folder_path = self.learnware_folder_list[learnware_id] diff --git a/learnware/market/easy/checker.py b/learnware/market/easy/checker.py index efb48b2..ce5d2b4 100644 --- a/learnware/market/easy/checker.py +++ b/learnware/market/easy/checker.py @@ -1,4 +1,6 @@ import traceback +import numpy as np +import torch from ..base import LearnwareChecker from ...logger import get_module_logger diff --git a/learnware/market/easy/database_ops.py b/learnware/market/easy/database_ops.py index 48ed0fb..61bc02d 100644 --- a/learnware/market/easy/database_ops.py +++ b/learnware/market/easy/database_ops.py @@ -1,10 +1,10 @@ from sqlalchemy.ext.declarative import declarative_base from sqlalchemy import create_engine, text -from sqlalchemy import Column, Text, String +from sqlalchemy import Column, Integer, Text, DateTime, String import os import json -from ...learnware import get_learnware_from_dirpath -from ...logger import get_module_logger +from ..learnware import get_learnware_from_dirpath +from ..logger import get_module_logger logger = get_module_logger("database") DeclarativeBase = declarative_base() @@ -117,17 +117,6 @@ class DatabaseOperations(object): pass pass - def update_learnware_semantic_spec(self, learnware_id: str, semantic_spec: dict): - with self.engine.connect() as conn: - semantic_spec_str = json.dumps(semantic_spec) - conn.execute( - text("UPDATE tb_learnware SET semantic_spec=:semantic_spec WHERE id=:id;"), - dict(id=learnware_id, semantic_spec=semantic_spec_str), - ) - conn.commit() - pass - pass - def delete_learnware(self, id: str): with self.engine.connect() as conn: conn.execute(text("DELETE FROM tb_learnware WHERE id=:id;"), dict(id=id)) @@ -146,6 +135,16 @@ class DatabaseOperations(object): pass pass + def update_learnware_use_flag(self, id: str, use_flag: str): + with self.engine.connect() as conn: + r = conn.execute( + text("UPDATE tb_learnware SET use_flag=:use_flag WHERE id=:id;"), + dict(id=id, use_flag=use_flag), + ) + conn.commit() + pass + pass + def load_market(self): with self.engine.connect() as conn: cursor = conn.execute(text("SELECT id, semantic_spec, zip_path, folder_path, use_flag FROM tb_learnware;")) @@ -153,6 +152,7 @@ class DatabaseOperations(object): learnware_list = {} zip_list = {} folder_list = {} + use_flags = {} max_count = 0 for id, semantic_spec, zip_path, folder_path, use_flag in cursor: @@ -166,10 +166,11 @@ class DatabaseOperations(object): # assert new_learnware is not None zip_list[id] = zip_path folder_list[id] = folder_path + use_flags[id] = use_flag max_count = max(max_count, int(id)) pass - return learnware_list, zip_list, folder_list, max_count + 1 + return learnware_list, zip_list, folder_list, use_flags, max_count + 1 pass pass diff --git a/learnware/market/easy/organizer.py b/learnware/market/easy/organizer.py index d3a70f8..7619e8a 100644 --- a/learnware/market/easy/organizer.py +++ b/learnware/market/easy/organizer.py @@ -4,6 +4,7 @@ import copy import torch import zipfile import traceback +import tempfile import numpy as np import pandas as pd from rapidfuzz import fuzz @@ -11,8 +12,10 @@ from cvxopt import solvers, matrix from shutil import copyfile, rmtree from typing import Tuple, Any, List, Union, Dict +from .database_ops import DatabaseOperations +from .checker import EasyChecker from ..base import LearnwareMarket, BaseUserInfo -from ..database_ops import DatabaseOperations + from ... import utils from ...config import C as conf @@ -28,8 +31,11 @@ logger = get_module_logger("easy_organizer") class EasyOrganizer(LearnwareOrganizer): - def reset(self, market_id, rebuild=False): - self.market_id = market_id + def __init__(self, market_id, checker: 'EasyChecker' = None, rebuild: bool = False): + self.reset(market_id=market_id, checker=checker, rebuild=rebuild) + + def reset(self, market_id, checker: EasyChecker = None, rebuild: bool = False): + super(EasyOrganizer, self).reset(market_id=market_id, checker=checker) self.reload_market(rebuild=rebuild) def reload_market(self, rebuild=False) -> bool: @@ -48,6 +54,7 @@ class EasyOrganizer(LearnwareOrganizer): self.learnware_list = {} # id: Learnware self.learnware_zip_list = {} self.learnware_folder_list = {} + self.use_flags = {} self.count = 0 self.semantic_spec_list = conf.semantic_specs self.dbops = DatabaseOperations(conf.database_url, "market_" + self.market_id) @@ -63,10 +70,10 @@ class EasyOrganizer(LearnwareOrganizer): os.makedirs(self.learnware_pool_path, exist_ok=True) os.makedirs(self.learnware_zip_pool_path, exist_ok=True) os.makedirs(self.learnware_folder_pool_path, exist_ok=True) - self.learnware_list, self.learnware_zip_list, self.learnware_folder_list, self.count = self.dbops.load_market() + self.learnware_list, self.learnware_zip_list, self.learnware_folder_list, self.use_flags, self.count = self.dbops.load_market() - def add_learnware(self, zip_path: str, semantic_spec: dict, learnware_id: str = None, check: bool = False) -> Tuple[str, bool]: + def add_learnware(self, zip_path: str, semantic_spec: dict, id: str = None, check: bool = True) -> Tuple[str, bool]: """Add a learnware into the market. .. note:: @@ -92,24 +99,24 @@ class EasyOrganizer(LearnwareOrganizer): if not os.path.exists(zip_path): logger.warning("Zip Path NOT Found! Fail to add learnware.") - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE try: if len(semantic_spec["Data"]["Values"]) == 0: logger.warning("Illegal semantic specification, please choose Data.") - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE if len(semantic_spec["Task"]["Values"]) == 0: logger.warning("Illegal semantic specification, please choose Task.") - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE if len(semantic_spec["Library"]["Values"]) == 0: logger.warning("Illegal semantic specification, please choose Device.") - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE if len(semantic_spec["Name"]["Values"]) == 0: logger.warning("Illegal semantic specification, please provide Name.") - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE if len(semantic_spec["Description"]["Values"]) == 0 and len(semantic_spec["Scenario"]["Values"]) == 0: logger.warning("Illegal semantic specification, please provide Scenario or Description.") - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE if ( semantic_spec["Data"]["Type"] != "Class" or semantic_spec["Task"]["Type"] != "Class" @@ -119,17 +126,15 @@ class EasyOrganizer(LearnwareOrganizer): or semantic_spec["Description"]["Type"] != "String" ): logger.warning("Illegal semantic specification, please provide the right type.") - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE except: print(semantic_spec) logger.warning("Illegal semantic specification, some keys are missing.") - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE logger.info("Get new learnware from %s" % (zip_path)) - if learnware_id is not None: - id = learnware_id - else: - id = "%08d" % (self.count) + + id = id if id is not None else "%08d" % (self.count) target_zip_dir = os.path.join(self.learnware_zip_pool_path, "%s.zip" % (id)) target_folder_dir = os.path.join(self.learnware_folder_pool_path, id) copyfile(zip_path, target_zip_dir) @@ -148,137 +153,167 @@ class EasyOrganizer(LearnwareOrganizer): rmtree(target_folder_dir) except: pass - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE if new_learnware is None: - return None, self.INVALID_LEARNWARE + return None, EasyChecker.INVALID_LEARNWARE - if check and self.checker + learnwere_status = EasyChecker.USABLE_LEARWARE if check is False else self.checker.check_learnware(new_learnware) self.dbops.add_learnware( id=id, semantic_spec=semantic_spec, zip_path=target_zip_dir, folder_path=target_folder_dir, - use_flag=LearnwareChecker.USABLE_LEARWARE, + use_flag=learnwere_status, ) self.learnware_list[id] = new_learnware self.learnware_zip_list[id] = target_zip_dir self.learnware_folder_list[id] = target_folder_dir + self.use_flags[id] = learnwere_status self.count += 1 - return id, LearnwareChecker.USABLE_LEARWARE - - def add_learnware(self, zip_path: str, semantic_spec: dict) -> Tuple[str, bool]: - """Add a learnware into the market. - - .. note:: - - Given a prediction of a certain time, all signals before this time will be prepared well. + return id, learnwere_status + def delete_learnware(self, id: str) -> bool: + """Delete Learnware from market Parameters ---------- - zip_path : str - Filepath for learnware model, a zipped file. - semantic_spec : dict - semantic_spec for new learnware, in dictionary format. + id : str + Learnware to be deleted Returns ------- - Tuple[str, int] - - str indicating model_id - - int indicating what the flag of learnware is added. - + bool + True for successful operation. + False for id not found. """ - semantic_spec = copy.deepcopy(semantic_spec) - - if not os.path.exists(zip_path): - logger.warning("Zip Path NOT Found! Fail to add learnware.") - return None, self.INVALID_LEARNWARE - - try: - if len(semantic_spec["Data"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please choose Data.") - return None, self.INVALID_LEARNWARE - if len(semantic_spec["Task"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please choose Task.") - return None, self.INVALID_LEARNWARE - if len(semantic_spec["Library"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please choose Device.") - return None, self.INVALID_LEARNWARE - if len(semantic_spec["Name"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please provide Name.") - return None, self.INVALID_LEARNWARE - if len(semantic_spec["Description"]["Values"]) == 0 and len(semantic_spec["Scenario"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please provide Scenario or Description.") - return None, self.INVALID_LEARNWARE - if ( - semantic_spec["Data"]["Type"] != "Class" - or semantic_spec["Task"]["Type"] != "Class" - or semantic_spec["Library"]["Type"] != "Class" - or semantic_spec["Scenario"]["Type"] != "Tag" - or semantic_spec["Name"]["Type"] != "String" - or semantic_spec["Description"]["Type"] != "String" - ): - logger.warning("Illegal semantic specification, please provide the right type.") - return None, self.INVALID_LEARNWARE - except: - logger.info(f"Semantic specification: {semantic_spec}") - logger.warning("Illegal semantic specification, some keys are missing.") - return None, self.INVALID_LEARNWARE - - logger.info("Get new learnware from %s" % (zip_path)) - id = "%08d" % (self.count) - target_zip_dir = os.path.join(self.learnware_zip_pool_path, "%s.zip" % (id)) - target_folder_dir = os.path.join(self.learnware_folder_pool_path, id) - copyfile(zip_path, target_zip_dir) - - with zipfile.ZipFile(target_zip_dir, "r") as z_file: - z_file.extractall(target_folder_dir) - logger.info("Learnware move to %s, and unzip to %s" % (target_zip_dir, target_folder_dir)) - - try: - new_learnware = get_learnware_from_dirpath( + if not id in self.learnware_list: + logger.warning("Learnware id:'{}' NOT Found!".format(id)) + return False + + zip_dir = self.learnware_zip_list[id] + os.remove(zip_dir) + folder_dir = self.learnware_folder_list[id] + rmtree(folder_dir) + self.learnware_list.pop(id) + self.learnware_zip_list.pop(id) + self.learnware_folder_list.pop(id) + self.use_flags.pop(id) + self.dbops.delete_learnware(id=id) + + return True + + def update_learnware( + self, id: str, zip_path: str = None, semantic_spec: dict = None, check: bool = True + ): + """TODO: update should pass the semantic check too + """ + assert zip_path is None and semantic_spec is None, f"at least one of 'zip_path' and 'semantic_spec' should not be None when update learnware" + + if semantic_spec is not None: + self.dbops.update_learnware_semantic_specification(id, semantic_spec) + else: + semantic_spec = self.learnware_list[id] + + if zip_path is not None: + target_zip_dir = self.learnware_zip_list[id] + target_folder_dir = self.learnware_folder_list[id] + + if check: + with tempfile.TemporaryDirectory(prefix="learnware_") as tempdir: + with zipfile.ZipFile(zip_path, "r") as z_file: + z_file.extractall(tempdir) + + try: + new_learnware = get_learnware_from_dirpath( + id=id, semantic_spec=semantic_spec, learnware_dirpath=tempdir + ) + except Exception: + return False, EasyChecker.INVALID_LEARNWARE + + if new_learnware is None: + return False, EasyChecker.INVALID_LEARNWARE + + learnwere_status = self.checker.check_learnware(new_learnware) + else: + learnwere_status = EasyChecker.USABLE_LEARWARE + + self.dbops.update_learnware_use_flag(id, learnwere_status) + copyfile(zip_path, target_zip_dir) + with zipfile.ZipFile(target_zip_dir, "r") as z_file: + z_file.extractall(target_folder_dir) + self.learnware_list[id] = get_learnware_from_dirpath( id=id, semantic_spec=semantic_spec, learnware_dirpath=target_folder_dir ) - except: - try: - os.remove(target_zip_dir) - rmtree(target_folder_dir) - except: - pass - return None, self.INVALID_LEARNWARE + + return True, learnwere_status + + else: + self.learnware_list[id].get_specification().update_semantic_spec(semantic_spec) + return self.use_flags[id] + + def get_learnware_by_ids(self, ids: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: + """Search learnware by id or list of ids. - if new_learnware is None: - return None, self.INVALID_LEARNWARE + Parameters + ---------- + ids : Union[str, List[str]] + Give a id or a list of ids + str: id of targer learware + List[str]: A list of ids of target learnwares - check_flag = self.check_learnware(new_learnware) + Returns + ------- + Union[Learnware, List[Learnware]] + Return target learnware or list of target learnwares. + None for Learnware NOT Found. + """ + if isinstance(ids, list): + ret = [] + for id in ids: + if id in self.learnware_list: + ret.append(self.learnware_list[id]) + else: + logger.warning("Learnware ID '%s' NOT Found!" % (id)) + ret.append(None) + return ret + else: + try: + return self.learnware_list[ids] + except: + logger.warning("Learnware ID '%s' NOT Found!" % (ids)) + return None - self.dbops.add_learnware( - id=id, - semantic_spec=semantic_spec, - zip_path=target_zip_dir, - folder_path=target_folder_dir, - use_flag=check_flag, - ) + def get_learnware_path_by_ids(self, ids: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: + """Get Zipped Learnware file by id - self.learnware_list[id] = new_learnware - self.learnware_zip_list[id] = target_zip_dir - self.learnware_folder_list[id] = target_folder_dir - self.count += 1 - return id, check_flag - + Parameters + ---------- + ids : Union[str, List[str]] + Give a id or a list of ids + str: id of targer learware + List[str]: A list of ids of target learnwares - def get_learnware_ids(self, top:int = None): - if top is None: - return list(self.learnware_list.keys()) - else: - return list(self.learnware_list.keys())[:top] - - - def get_learnwares(self, top:int = None): - if top is None: - return list(self.learnware_list.values()) + Returns + ------- + Union[Learnware, List[Learnware]] + Return the path for target learnware or list of path. + None for Learnware NOT Found. + """ + if isinstance(ids, list): + ret = [] + for id in ids: + if id in self.learnware_zip_list: + ret.append(self.learnware_zip_list[id]) + else: + logger.warning("Learnware ID '%s' NOT Found!" % (id)) + ret.append(None) + return ret else: - return list(self.learnware_list.values())[:top] \ No newline at end of file + try: + return self.learnware_zip_list[ids] + except: + logger.warning("Learnware ID '%s' NOT Found!" % (ids)) + return None \ No newline at end of file From 21241f22e66da120c6fec04d509534414d6328d3 Mon Sep 17 00:00:00 2001 From: bxdd Date: Thu, 26 Oct 2023 22:33:04 +0800 Subject: [PATCH 08/35] [FIX] fix bugs --- learnware/market/base.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/learnware/market/base.py b/learnware/market/base.py index 2bec206..84a52b4 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -252,7 +252,7 @@ class LearnwareSearcher: def reset(self, organizer): self.learnware_oganizer = organizer - def __call__(self, user_info: BaseUserInfo) + def __call__(self, user_info: BaseUserInfo): """Search learnwares based on user_info Parameters From 3e79d808a2d34a348921ee529d4999b90a463c87 Mon Sep 17 00:00:00 2001 From: bxdd Date: Fri, 27 Oct 2023 02:01:28 +0800 Subject: [PATCH 09/35] [MNT] del searcher file --- learnware/market/{easy => easymarket}/__init__.py | 0 learnware/market/{easy => easymarket}/checker.py | 0 learnware/market/{easy => easymarket}/database_ops.py | 0 learnware/market/{easy => easymarket}/organizer.py | 0 learnware/market/{easy => easymarket}/searcher.py | 0 learnware/market/searcher.py | 10 ---------- 6 files changed, 10 deletions(-) rename learnware/market/{easy => easymarket}/__init__.py (100%) rename learnware/market/{easy => easymarket}/checker.py (100%) rename learnware/market/{easy => easymarket}/database_ops.py (100%) rename learnware/market/{easy => easymarket}/organizer.py (100%) rename learnware/market/{easy => easymarket}/searcher.py (100%) delete mode 100644 learnware/market/searcher.py diff --git a/learnware/market/easy/__init__.py b/learnware/market/easymarket/__init__.py similarity index 100% rename from learnware/market/easy/__init__.py rename to learnware/market/easymarket/__init__.py diff --git a/learnware/market/easy/checker.py b/learnware/market/easymarket/checker.py similarity index 100% rename from learnware/market/easy/checker.py rename to learnware/market/easymarket/checker.py diff --git a/learnware/market/easy/database_ops.py b/learnware/market/easymarket/database_ops.py similarity index 100% rename from learnware/market/easy/database_ops.py rename to learnware/market/easymarket/database_ops.py diff --git a/learnware/market/easy/organizer.py b/learnware/market/easymarket/organizer.py similarity index 100% rename from learnware/market/easy/organizer.py rename to learnware/market/easymarket/organizer.py diff --git a/learnware/market/easy/searcher.py b/learnware/market/easymarket/searcher.py similarity index 100% rename from learnware/market/easy/searcher.py rename to learnware/market/easymarket/searcher.py diff --git a/learnware/market/searcher.py b/learnware/market/searcher.py deleted file mode 100644 index 09e6f30..0000000 --- a/learnware/market/searcher.py +++ /dev/null @@ -1,10 +0,0 @@ - - -from typing import Tuple, Any, List - -from .base import BaseUserInfo -from ..learnware import Learnware -from ..logger import get_module_logger - -logger = get_module_logger('model') - From 0437f09f60fb73e4d1f0082cf89f6958a6e0fe32 Mon Sep 17 00:00:00 2001 From: bxdd Date: Fri, 27 Oct 2023 02:19:08 +0800 Subject: [PATCH 10/35] [MNT] modify update_learnware in organizer --- learnware/market/easymarket/organizer.py | 54 ++- learnware/market/easymarket/searcher.py | 539 ++++++++++++++++++++++- 2 files changed, 578 insertions(+), 15 deletions(-) diff --git a/learnware/market/easymarket/organizer.py b/learnware/market/easymarket/organizer.py index 7619e8a..ecbea5b 100644 --- a/learnware/market/easymarket/organizer.py +++ b/learnware/market/easymarket/organizer.py @@ -73,7 +73,7 @@ class EasyOrganizer(LearnwareOrganizer): self.learnware_list, self.learnware_zip_list, self.learnware_folder_list, self.use_flags, self.count = self.dbops.load_market() - def add_learnware(self, zip_path: str, semantic_spec: dict, id: str = None, check: bool = True) -> Tuple[str, bool]: + def add_learnware(self, zip_path: str, semantic_spec: dict, id: str = None, check_status: int = None) -> Tuple[str, bool]: """Add a learnware into the market. .. note:: @@ -158,7 +158,7 @@ class EasyOrganizer(LearnwareOrganizer): if new_learnware is None: return None, EasyChecker.INVALID_LEARNWARE - learnwere_status = EasyChecker.USABLE_LEARWARE if check is False else self.checker.check_learnware(new_learnware) + learnwere_status = check_status if check_status is not None else self.checker.check_learnware(new_learnware) self.dbops.add_learnware( id=id, @@ -206,11 +206,29 @@ class EasyOrganizer(LearnwareOrganizer): return True def update_learnware( - self, id: str, zip_path: str = None, semantic_spec: dict = None, check: bool = True + self, id: str, zip_path: str = None, semantic_spec: dict = None, check_status: int = None ): - """TODO: update should pass the semantic check too + """update learnware with zip_path and semantic_specification + TODO: update should pass the semantic check too + + Parameters + ---------- + id : str + _description_ + zip_path : str, optional + _description_, by default None + semantic_spec : dict, optional + _description_, by default None + check_status : int, optional + _description_, by default None + + Returns + ------- + _type_ + _description_ """ assert zip_path is None and semantic_spec is None, f"at least one of 'zip_path' and 'semantic_spec' should not be None when update learnware" + assert check_status != EasyChecker.INVALID_LEARNWARE, f"'check_status' can not be INVALID_LEARNWARE" if semantic_spec is not None: self.dbops.update_learnware_semantic_specification(id, semantic_spec) @@ -221,7 +239,7 @@ class EasyOrganizer(LearnwareOrganizer): target_zip_dir = self.learnware_zip_list[id] target_folder_dir = self.learnware_folder_list[id] - if check: + if check_status is None: with tempfile.TemporaryDirectory(prefix="learnware_") as tempdir: with zipfile.ZipFile(zip_path, "r") as z_file: z_file.extractall(tempdir) @@ -231,14 +249,14 @@ class EasyOrganizer(LearnwareOrganizer): id=id, semantic_spec=semantic_spec, learnware_dirpath=tempdir ) except Exception: - return False, EasyChecker.INVALID_LEARNWARE + return EasyChecker.INVALID_LEARNWARE if new_learnware is None: - return False, EasyChecker.INVALID_LEARNWARE + return EasyChecker.INVALID_LEARNWARE learnwere_status = self.checker.check_learnware(new_learnware) else: - learnwere_status = EasyChecker.USABLE_LEARWARE + learnwere_status = check_status self.dbops.update_learnware_use_flag(id, learnwere_status) copyfile(zip_path, target_zip_dir) @@ -247,9 +265,7 @@ class EasyOrganizer(LearnwareOrganizer): self.learnware_list[id] = get_learnware_from_dirpath( id=id, semantic_spec=semantic_spec, learnware_dirpath=target_folder_dir ) - - return True, learnwere_status - + return learnwere_status else: self.learnware_list[id].get_specification().update_semantic_spec(semantic_spec) return self.use_flags[id] @@ -316,4 +332,18 @@ class EasyOrganizer(LearnwareOrganizer): return self.learnware_zip_list[ids] except: logger.warning("Learnware ID '%s' NOT Found!" % (ids)) - return None \ No newline at end of file + return None + + def get_learnware_ids(self, top:int = None) -> List[str]: + if top is None: + return list(self.learnware_list.keys()) + else: + return list(self.learnware_list.keys())[:top] + + + def get_learnwares(self, top:int = None) -> List[str]: + if top is None: + return list(self.learnware_list.values()) + else: + return list(self.learnware_list.values())[:top] + \ No newline at end of file diff --git a/learnware/market/easymarket/searcher.py b/learnware/market/easymarket/searcher.py index 427f51c..e5adcd8 100644 --- a/learnware/market/easymarket/searcher.py +++ b/learnware/market/easymarket/searcher.py @@ -1,14 +1,547 @@ +import torch +import numpy as np +from rapidfuzz import fuzz +from cvxopt import solvers, matrix from typing import Tuple, List -from ..base import LearnwareSearcher -from ...logger import get_module_logger +from ..base import BaseUserInfo, LearnwareSearcher from ...learnware import Learnware -from ...market import BaseUserInfo +from ...specification import RKMEStatSpecification +from ...logger import get_module_logger logger = get_module_logger('easy_seacher') class EasySearcher(LearnwareSearcher): + def _convert_dist_to_score( + self, dist_list: List[float], dist_epsilon: float = 0.01, min_score: float = 0.92 + ) -> List[float]: + """Convert mmd dist list into min_max score list + + Parameters + ---------- + dist_list : List[float] + The list of mmd distances from learnware rkmes to user rkme + dist_epsilon: float + The paramter for converting mmd dist to score + min_score: float + The minimum score for maximum returned score + + Returns + ------- + List[float] + The list of min_max scores of each learnware + """ + if len(dist_list) == 0: + return [] + + min_dist, max_dist = min(dist_list), max(dist_list) + if min_dist == max_dist: + return [1 for dist in dist_list] + else: + max_score = (max_dist - min_dist) / (max_dist - dist_epsilon) + + if min_dist < dist_epsilon: + dist_epsilon = min_dist + elif max_score < min_score: + dist_epsilon = max_dist - (max_dist - min_dist) / min_score + + return [(max_dist - dist) / (max_dist - dist_epsilon) for dist in dist_list] + + def _calculate_rkme_spec_mixture_weight( + self, + learnware_list: List[Learnware], + user_rkme: RKMEStatSpecification, + intermediate_K: np.ndarray = None, + intermediate_C: np.ndarray = None, + ) -> Tuple[List[float], float]: + """Calculate mixture weight for the learnware_list based on a user's rkme + + Parameters + ---------- + learnware_list : List[Learnware] + A list of existing learnwares + user_rkme : RKMEStatSpecification + User RKME statistical specification + intermediate_K : np.ndarray, optional + Intermediate kernel matrix K, by default None + intermediate_C : np.ndarray, optional + Intermediate inner product vector C, by default None + + Returns + ------- + Tuple[List[float], float] + The first is the list of mixture weights + The second is the mmd dist between the mixture of learnware rkmes and the user's rkme + """ + learnware_num = len(learnware_list) + RKME_list = [ + learnware.specification.get_stat_spec_by_name("RKMEStatSpecification") for learnware in learnware_list + ] + + if type(intermediate_K) == np.ndarray: + K = intermediate_K + else: + K = np.zeros((learnware_num, learnware_num)) + for i in range(K.shape[0]): + K[i, i] = RKME_list[i].inner_prod(RKME_list[i]) + for j in range(i + 1, K.shape[0]): + K[i, j] = K[j, i] = RKME_list[i].inner_prod(RKME_list[j]) + + if type(intermediate_C) == np.ndarray: + C = intermediate_C + else: + C = np.zeros((learnware_num, 1)) + for i in range(C.shape[0]): + C[i, 0] = user_rkme.inner_prod(RKME_list[i]) + + K = torch.from_numpy(K).double().to(user_rkme.device) + C = torch.from_numpy(C).double().to(user_rkme.device) + + # beta can be negative + # weight = torch.linalg.inv(K + torch.eye(K.shape[0]).to(user_rkme.device) * 1e-5) @ C + + # beta must be nonnegative + n = K.shape[0] + P = matrix(K.cpu().numpy()) + q = matrix(-C.cpu().numpy()) + G = matrix(-np.eye(n)) + h = matrix(np.zeros((n, 1))) + A = matrix(np.ones((1, n))) + b = matrix(np.ones((1, 1))) + solvers.options["show_progress"] = False + sol = solvers.qp(P, q, G, h, A, b) + weight = np.array(sol["x"]) + weight = torch.from_numpy(weight).reshape(-1).double().to(user_rkme.device) + score = user_rkme.inner_prod(user_rkme) + 2 * sol["primal objective"] + + return weight.detach().cpu().numpy().reshape(-1), score + + def _calculate_intermediate_K_and_C( + self, + learnware_list: List[Learnware], + user_rkme: RKMEStatSpecification, + intermediate_K: np.ndarray = None, + intermediate_C: np.ndarray = None, + ) -> Tuple[np.ndarray, np.ndarray]: + """Incrementally update the values of intermediate_K and intermediate_C + + Parameters + ---------- + learnware_list : List[Learnware] + The list of learnwares up till now + user_rkme : RKMEStatSpecification + User RKME statistical specification + intermediate_K : np.ndarray, optional + Intermediate kernel matrix K, by default None + intermediate_C : np.ndarray, optional + Intermediate inner product vector C, by default None + + Returns + ------- + Tuple[np.ndarray, np.ndarray] + The first is the intermediate value of K + The second is the intermediate value of C + """ + num = intermediate_K.shape[0] - 1 + RKME_list = [ + learnware.specification.get_stat_spec_by_name("RKMEStatSpecification") for learnware in learnware_list + ] + for i in range(intermediate_K.shape[0]): + intermediate_K[num, i] = RKME_list[-1].inner_prod(RKME_list[i]) + intermediate_C[num, 0] = user_rkme.inner_prod(RKME_list[-1]) + return intermediate_K, intermediate_C + + def _search_by_rkme_spec_mixture_auto( + self, + learnware_list: List[Learnware], + user_rkme: RKMEStatSpecification, + max_search_num: int, + weight_cutoff: float = 0.98, + ) -> Tuple[float, List[float], List[Learnware]]: + """Select learnwares based on a total mixture ratio, then recalculate their mixture weights + + Parameters + ---------- + learnware_list : List[Learnware] + The list of learnwares whose mixture approximates the user's rkme + user_rkme : RKMEStatSpecification + User RKME statistical specification + max_search_num : int + The maximum number of the returned learnwares + weight_cutoff : float, optional + The ratio for selecting out the mose relevant learnwares, by default 0.9 + + Returns + ------- + Tuple[float, List[float], List[Learnware]] + The first is the mixture mmd dist + The second is the list of weight + The third is the list of Learnware + """ + learnware_num = len(learnware_list) + if learnware_num == 0: + return [], [] + if learnware_num < max_search_num: + logger.warning("Available Learnware num less than search_num!") + max_search_num = learnware_num + + weight, _ = self._calculate_rkme_spec_mixture_weight(learnware_list, user_rkme) + sort_by_weight_idx_list = sorted(range(learnware_num), key=lambda k: weight[k], reverse=True) + + weight_sum = 0 + mixture_list = [] + for idx in sort_by_weight_idx_list: + weight_sum += weight[idx] + if weight_sum <= weight_cutoff: + mixture_list.append(learnware_list[idx]) + else: + break + + if len(mixture_list) <= 1: + mixture_list = [learnware_list[sort_by_weight_idx_list[0]]] + mixture_weight = [1] + mmd_dist = user_rkme.dist(mixture_list[0].specification.get_stat_spec_by_name("RKMEStatSpecification")) + else: + if len(mixture_list) > max_search_num: + mixture_list = mixture_list[:max_search_num] + mixture_weight, mmd_dist = self._calculate_rkme_spec_mixture_weight(mixture_list, user_rkme) + + return mmd_dist, mixture_weight, mixture_list + + def _filter_by_rkme_spec_single( + self, + sorted_score_list: List[float], + learnware_list: List[Learnware], + filter_score: float = 0.5, + min_num: int = 15, + ) -> Tuple[List[float], List[Learnware]]: + """Filter search result of _search_by_rkme_spec_single + + Parameters + ---------- + sorted_score_list : List[float] + The list of score transformed by mmd dist + learnware_list : List[Learnware] + The list of learnwares whose mixture approximates the user's rkme + filter_score: float + The learnware whose score is lower than filter_score will be filtered + min_num: int + The minimum number of returned learnwares + + Returns + ------- + Tuple[List[float], List[Learnware]] + the first is the list of score + the second is the list of Learnware + """ + idx = min(min_num, len(learnware_list)) + while idx < len(learnware_list): + if sorted_score_list[idx] < filter_score: + break + idx = idx + 1 + return sorted_score_list[:idx], learnware_list[:idx] + + def _filter_by_rkme_spec_dimension( + self, learnware_list: List[Learnware], user_rkme: RKMEStatSpecification + ) -> List[Learnware]: + """Filter learnwares whose rkme dimension different from user_rkme + + Parameters + ---------- + learnware_list : List[Learnware] + The list of learnwares whose mixture approximates the user's rkme + user_rkme : RKMEStatSpecification + User RKME statistical specification + + Returns + ------- + List[Learnware] + Learnwares whose rkme dimensions equal user_rkme in user_info + """ + filtered_learnware_list = [] + user_rkme_dim = str(list(user_rkme.get_z().shape)[1:]) + + for learnware in learnware_list: + rkme = learnware.specification.get_stat_spec_by_name("RKMEStatSpecification") + rkme_dim = str(list(rkme.get_z().shape)[1:]) + if rkme_dim == user_rkme_dim: + filtered_learnware_list.append(learnware) + + return filtered_learnware_list + + def _search_by_rkme_spec_mixture_greedy( + self, + learnware_list: List[Learnware], + user_rkme: RKMEStatSpecification, + max_search_num: int, + score_cutoff: float = 0.001, + ) -> Tuple[float, List[float], List[Learnware]]: + """Greedily match learnwares such that their mixture become closer and closer to user's rkme + + Parameters + ---------- + learnware_list : List[Learnware] + The list of learnwares whose mixture approximates the user's rkme + user_rkme : RKMEStatSpecification + User RKME statistical specification + max_search_num : int + The maximum number of the returned learnwares + score_cutof: float + The minimum mmd dist as threshold to stop further rkme_spec matching + + Returns + ------- + Tuple[float, List[float], List[Learnware]] + The first is the mixture mmd dist + The second is the list of weight + The third is the list of Learnware + """ + learnware_num = len(learnware_list) + if learnware_num == 0: + return None, [], [] + if learnware_num < max_search_num: + logger.warning("Available Learnware num less than search_num!") + max_search_num = learnware_num + + flag_list = [0 for _ in range(learnware_num)] + mixture_list, mmd_dist = [], None + intermediate_K, intermediate_C = np.zeros((1, 1)), np.zeros((1, 1)) + + for k in range(max_search_num): + idx_min, score_min = -1, -1 + weight_min = None + mixture_list.append(None) + + if k != 0: + intermediate_K = np.c_[intermediate_K, np.zeros((k, 1))] + intermediate_K = np.r_[intermediate_K, np.zeros((1, k + 1))] + intermediate_C = np.r_[intermediate_C, np.zeros((1, 1))] + + for idx in range(len(learnware_list)): + if flag_list[idx] == 0: + mixture_list[-1] = learnware_list[idx] + intermediate_K, intermediate_C = self._calculate_intermediate_K_and_C( + mixture_list, user_rkme, intermediate_K, intermediate_C + ) + weight, score = self._calculate_rkme_spec_mixture_weight( + mixture_list, user_rkme, intermediate_K, intermediate_C + ) + if idx_min == -1 or score < score_min: + idx_min, score_min, weight_min = idx, score, weight + + mmd_dist = score_min + mixture_list[-1] = learnware_list[idx_min] + if score_min < score_cutoff: + break + else: + flag_list[idx_min] = 1 + intermediate_K, intermediate_C = self._calculate_intermediate_K_and_C( + mixture_list, user_rkme, intermediate_K, intermediate_C + ) + + return mmd_dist, weight_min, mixture_list + + def _search_by_rkme_spec_single( + self, learnware_list: List[Learnware], user_rkme: RKMEStatSpecification + ) -> Tuple[List[float], List[Learnware]]: + """Calculate the distances between learnwares in the given learnware_list and user_rkme + + Parameters + ---------- + learnware_list : List[Learnware] + The list of learnwares whose mixture approximates the user's rkme + user_rkme : RKMEStatSpecification + user RKME statistical specification + + Returns + ------- + Tuple[List[float], List[Learnware]] + the first is the list of mmd dist + the second is the list of Learnware + both lists are sorted by mmd dist + """ + RKME_list = [ + learnware.specification.get_stat_spec_by_name("RKMEStatSpecification") for learnware in learnware_list + ] + mmd_dist_list = [] + for RKME in RKME_list: + mmd_dist = RKME.dist(user_rkme) + mmd_dist_list.append(mmd_dist) + + sorted_idx_list = sorted(range(len(learnware_list)), key=lambda k: mmd_dist_list[k]) + sorted_dist_list = [mmd_dist_list[idx] for idx in sorted_idx_list] + sorted_learnware_list = [learnware_list[idx] for idx in sorted_idx_list] + + return sorted_dist_list, sorted_learnware_list + + def _search_by_semantic_spec_exact( + self, learnware_list: List[Learnware], user_info: BaseUserInfo + ) -> List[Learnware]: + def match_semantic_spec(semantic_spec1, semantic_spec2): + """ + semantic_spec1: semantic spec input by user + semantic_spec2: semantic spec in database + """ + if semantic_spec1.keys() != semantic_spec2.keys(): + # sematic spec in database may contain more keys than user input + pass + + name2 = semantic_spec2["Name"]["Values"].lower() + description2 = semantic_spec2["Description"]["Values"].lower() + + for key in semantic_spec1.keys(): + v1 = semantic_spec1[key]["Values"] + v2 = semantic_spec2[key]["Values"] + + if len(v1) == 0: + # user input is empty, no need to search + continue + + if key in ("Name", "Description"): + v1 = v1.lower() + if v1 not in name2 and v1 not in description2: + return False + pass + else: + if len(v2) == 0: + # user input contains some key that is not in database + return False + + if semantic_spec1[key]["Type"] == "Class": + if isinstance(v1, list): + v1 = v1[0] + if isinstance(v2, list): + v2 = v2[0] + if v1 != v2: + return False + elif semantic_spec1[key]["Type"] == "Tag": + if not (set(v1) & set(v2)): + return False + pass + pass + pass + + return True + + match_learnwares = [] + for learnware in learnware_list: + learnware_semantic_spec = learnware.get_specification().get_semantic_spec() + user_semantic_spec = user_info.get_semantic_spec() + if match_semantic_spec(user_semantic_spec, learnware_semantic_spec): + match_learnwares.append(learnware) + logger.info("semantic_spec search: choose %d from %d learnwares" % (len(match_learnwares), len(learnware_list))) + return match_learnwares + + def _search_by_semantic_spec_fuzz( + self, learnware_list: List[Learnware], user_info: BaseUserInfo, max_num: int = 50000, min_score: float = 75.0 + ) -> List[Learnware]: + """Search learnware by fuzzy matching of semantic spec + + Parameters + ---------- + learnware_list : List[Learnware] + The list of learnwares + user_info : BaseUserInfo + user_info contains semantic_spec + max_num : int, optional + maximum number of learnwares returned, by default 50000 + min_score : float, optional + Minimum fuzzy matching score of learnwares returned, by default 30.0 + + Returns + ------- + List[Learnware] + The list of returned learnwares + """ + def _match_semantic_spec_tag(semantic_spec1, semantic_spec2) -> bool: + """Judge if tags of two semantic specs are consistent + + Parameters + ---------- + semantic_spec1 : + semantic spec input by user + semantic_spec2 : + semantic spec in database + + Returns + ------- + bool + consistent (True) or not consistent (False) + """ + for key in semantic_spec1.keys(): + v1 = semantic_spec1[key]["Values"] + v2 = semantic_spec2[key]["Values"] + + if len(v1) == 0: + # user input is empty, no need to search + continue + + if key not in "Name": + if len(v2) == 0: + # user input contains some key that is not in database + return False + + if semantic_spec1[key]["Type"] == "Class": + if isinstance(v1, list): + v1 = v1[0] + if isinstance(v2, list): + v2 = v2[0] + if v1 != v2: + return False + elif semantic_spec1[key]["Type"] == "Tag": + if not (set(v1) & set(v2)): + return False + return True + + matched_learnware_tag = [] + final_result = [] + user_semantic_spec = user_info.get_semantic_spec() + + for learnware in learnware_list: + learnware_semantic_spec = learnware.get_specification().get_semantic_spec() + if _match_semantic_spec_tag(user_semantic_spec, learnware_semantic_spec): + matched_learnware_tag.append(learnware) + + if len(matched_learnware_tag) > 0: + if "Name" in user_semantic_spec: + name_user = user_semantic_spec["Name"]["Values"].lower() + if len(name_user) > 0: + # Exact search + name_list = [learnware.get_specification().get_semantic_spec()["Name"]["Values"].lower() for learnware in matched_learnware_tag] + des_list = [learnware.get_specification().get_semantic_spec()["Description"]["Values"].lower() for learnware in matched_learnware_tag] + + matched_learnware_exact = [] + for i in range(len(name_list)): + if name_user in name_list[i] or name_user in des_list[i]: + matched_learnware_exact.append(matched_learnware_tag[i]) + + if len(matched_learnware_exact) == 0: + # Fuzzy search + matched_learnware_fuzz, fuzz_scores = [], [] + for i in range(len(name_list)): + score_name = fuzz.partial_ratio(name_user, name_list[i]) + score_des = fuzz.partial_ratio(name_user, des_list[i]) + final_score = max(score_name, score_des) + if final_score >= min_score: + matched_learnware_fuzz.append(matched_learnware_tag[i]) + fuzz_scores.append(final_score) + + # Sort by score + sort_idx = sorted(list(range(len(fuzz_scores))), key=lambda k: fuzz_scores[k], reverse=True)[:max_num] + final_result = [matched_learnware_fuzz[idx] for idx in sort_idx] + else: + final_result = matched_learnware_exact + else: + final_result = matched_learnware_tag + else: + final_result = matched_learnware_tag + + logger.info( + "semantic_spec search: choose %d from %d learnwares" % (len(final_result), len(learnware_list)) + ) + return final_result + def __call__(self, user_info: BaseUserInfo, max_search_num: int = 5, search_method: str = "greedy") -> Tuple[List[float], List[Learnware], float, List[Learnware]]: """Search learnwares based on user_info From c4a5b9c572819620c0c668e619952ae165ce8ddb Mon Sep 17 00:00:00 2001 From: bxdd Date: Fri, 27 Oct 2023 02:30:21 +0800 Subject: [PATCH 11/35] [MNT] update update_learnware in organizer --- learnware/market/easymarket/organizer.py | 74 ++++++++++++------------ 1 file changed, 38 insertions(+), 36 deletions(-) diff --git a/learnware/market/easymarket/organizer.py b/learnware/market/easymarket/organizer.py index ecbea5b..2295f63 100644 --- a/learnware/market/easymarket/organizer.py +++ b/learnware/market/easymarket/organizer.py @@ -229,46 +229,48 @@ class EasyOrganizer(LearnwareOrganizer): """ assert zip_path is None and semantic_spec is None, f"at least one of 'zip_path' and 'semantic_spec' should not be None when update learnware" assert check_status != EasyChecker.INVALID_LEARNWARE, f"'check_status' can not be INVALID_LEARNWARE" + + if zip_path is None and check_status is not None: + logger.warning("check_status will be ignored when zip_path is None for learnware update") + + learnware_zippath = self.learnware_zip_list[id] if zip_path is None else zip_path + semantic_spec = self.learnware_list[id].get_specification().get_semantic_spec() if semantic_spec is None else semantic_spec + + self.dbops.update_learnware_semantic_specification(id, semantic_spec) - if semantic_spec is not None: - self.dbops.update_learnware_semantic_specification(id, semantic_spec) + target_zip_dir = self.learnware_zip_list[id] + target_folder_dir = self.learnware_folder_list[id] + + if check_status is None and zip_path is not None: + with tempfile.TemporaryDirectory(prefix="learnware_") as tempdir: + with zipfile.ZipFile(zip_path, "r") as z_file: + z_file.extractall(tempdir) + + try: + new_learnware = get_learnware_from_dirpath( + id=id, semantic_spec=semantic_spec, learnware_dirpath=tempdir + ) + except Exception: + return EasyChecker.INVALID_LEARNWARE + + if new_learnware is None: + return EasyChecker.INVALID_LEARNWARE + + learnwere_status = self.checker.check_learnware(new_learnware) else: - semantic_spec = self.learnware_list[id] + learnwere_status = self.use_flags[id] if zip_path is None else check_status - if zip_path is not None: - target_zip_dir = self.learnware_zip_list[id] - target_folder_dir = self.learnware_folder_list[id] - - if check_status is None: - with tempfile.TemporaryDirectory(prefix="learnware_") as tempdir: - with zipfile.ZipFile(zip_path, "r") as z_file: - z_file.extractall(tempdir) - - try: - new_learnware = get_learnware_from_dirpath( - id=id, semantic_spec=semantic_spec, learnware_dirpath=tempdir - ) - except Exception: - return EasyChecker.INVALID_LEARNWARE - - if new_learnware is None: - return EasyChecker.INVALID_LEARNWARE + copyfile(zip_path, target_zip_dir) + with zipfile.ZipFile(target_zip_dir, "r") as z_file: + z_file.extractall(target_folder_dir) - learnwere_status = self.checker.check_learnware(new_learnware) - else: - learnwere_status = check_status - - self.dbops.update_learnware_use_flag(id, learnwere_status) - copyfile(zip_path, target_zip_dir) - with zipfile.ZipFile(target_zip_dir, "r") as z_file: - z_file.extractall(target_folder_dir) - self.learnware_list[id] = get_learnware_from_dirpath( - id=id, semantic_spec=semantic_spec, learnware_dirpath=target_folder_dir - ) - return learnwere_status - else: - self.learnware_list[id].get_specification().update_semantic_spec(semantic_spec) - return self.use_flags[id] + self.learnware_list[id] = get_learnware_from_dirpath( + id=id, semantic_spec=semantic_spec, learnware_dirpath=target_folder_dir + ) + self.use_flags[id] = learnwere_status + self.dbops.update_learnware_use_flag(id, learnwere_status) + return learnwere_status + def get_learnware_by_ids(self, ids: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """Search learnware by id or list of ids. From 25af8c6ad4004fccc4d74d1084e25588d7f6ad16 Mon Sep 17 00:00:00 2001 From: bxdd Date: Fri, 27 Oct 2023 03:07:37 +0800 Subject: [PATCH 12/35] [ENH] add organizer for anchor, evolve marker --- learnware/market/anchor/__init__.py | 37 +++++++++ .../market/{anchor.py => anchor/organizer.py} | 78 ++++++------------- learnware/market/base.py | 3 - .../market/{easymarket => easy2}/__init__.py | 0 .../market/{easymarket => easy2}/checker.py | 0 .../{easymarket => easy2}/database_ops.py | 0 .../market/{easymarket => easy2}/organizer.py | 0 .../market/{easymarket => easy2}/searcher.py | 0 learnware/market/evolve/__init__.py | 0 .../market/{evolve.py => evolve/organizer.py} | 24 +++--- learnware/market/evolve_anchor/__init__.py | 0 .../organizer.py} | 22 +++--- learnware/market/hetergeneous/__init__.py | 0 .../organizer.py} | 17 ++-- 14 files changed, 87 insertions(+), 94 deletions(-) create mode 100644 learnware/market/anchor/__init__.py rename learnware/market/{anchor.py => anchor/organizer.py} (53%) rename learnware/market/{easymarket => easy2}/__init__.py (100%) rename learnware/market/{easymarket => easy2}/checker.py (100%) rename learnware/market/{easymarket => easy2}/database_ops.py (100%) rename learnware/market/{easymarket => easy2}/organizer.py (100%) rename learnware/market/{easymarket => easy2}/searcher.py (100%) create mode 100644 learnware/market/evolve/__init__.py rename learnware/market/{evolve.py => evolve/organizer.py} (65%) create mode 100644 learnware/market/evolve_anchor/__init__.py rename learnware/market/{evolve_anchor.py => evolve_anchor/organizer.py} (56%) create mode 100644 learnware/market/hetergeneous/__init__.py rename learnware/market/{heterogeneous_feature.py => hetergeneous/organizer.py} (85%) diff --git a/learnware/market/anchor/__init__.py b/learnware/market/anchor/__init__.py new file mode 100644 index 0000000..80680a9 --- /dev/null +++ b/learnware/market/anchor/__init__.py @@ -0,0 +1,37 @@ + + +class AnchoredUserInfo(BaseUserInfo): + """ + User Information for searching learnware (add the anchor design) + + - UserInfo contains the anchor list acquired from the market + - UserInfo can update stat_info based on anchors + """ + + def __init__(self, id: str, semantic_spec: dict = dict(), stat_info: dict = dict()): + super(AnchoredUserInfo, self).__init__(id, semantic_spec, stat_info) + self.anchor_learnware_list = {} # id: Learnware + + def add_anchor_learnware(self, learnware_id: str, learnware: Learnware): + """Add the anchor learnware acquired from the market + + Parameters + ---------- + learnware_id : str + Id of anchor learnware + learnware : Learnware + Anchor learnware for capturing user requirements + """ + self.anchor_learnware_list[learnware_id] = learnware + + def update_stat_info(self, name: str, item: Any): + """Update stat_info based on anchor learnwares + + Parameters + ---------- + name : str + Name of stat_info + item : Any + Statistical information calculated on anchor learnwares + """ + self.stat_info[name] = item diff --git a/learnware/market/anchor.py b/learnware/market/anchor/organizer.py similarity index 53% rename from learnware/market/anchor.py rename to learnware/market/anchor/organizer.py index bd912f3..e1213fc 100644 --- a/learnware/market/anchor.py +++ b/learnware/market/anchor/organizer.py @@ -1,9 +1,12 @@ -import os -from typing import Tuple, Any, List, Union, Dict +from typing import List, Dict, Tuple, Any -from ..learnware import Learnware -from .base import LearnwareMarket, BaseUserInfo +from ..base import BaseUserInfo +from ..easy2.organizer import EasyOrganizer +from ...logger import get_module_logger +from ...learnware import Learnware +from ...specification import BaseStatSpecification +logger = get_module_logger("evolve_organizer") class AnchoredUserInfo(BaseUserInfo): """ @@ -13,46 +16,29 @@ class AnchoredUserInfo(BaseUserInfo): - UserInfo can update stat_info based on anchors """ - def __init__(self, id: str, semantic_spec: dict = dict(), stat_info: dict = dict()): + def __init__(self, id: str, semantic_spec: dict = None, stat_info: dict = None, anchor_scores: dict = None): super(AnchoredUserInfo, self).__init__(id, semantic_spec, stat_info) - self.anchor_learnware_list = {} # id: Learnware + self.anchor_scores = {} if anchor_scores is None else anchor_scores - def add_anchor_learnware(self, learnware_id: str, learnware: Learnware): - """Add the anchor learnware acquired from the market + def update_anchor_score(self, id: str, score): + """Update score of anchor learnwares Parameters ---------- - learnware_id : str - Id of anchor learnware - learnware : Learnware - Anchor learnware for capturing user requirements - """ - self.anchor_learnware_list[learnware_id] = learnware - - def update_stat_info(self, name: str, item: Any): - """Update stat_info based on anchor learnwares - - Parameters - ---------- - name : str - Name of stat_info - item : Any - Statistical information calculated on anchor learnwares + id : str + id of anchor learnwares + score : Any + score of anchor learnwares """ - self.stat_info[name] = item - + self.anchor_scores[id] = score -class AnchoredMarket(LearnwareMarket): - """Add the anchor design to the LearnwareMarket - Parameters - ---------- - LearnwareMarket : _type_ - Basic market version +class AnchoredOrganizer(EasyOrganizer): + """Organize learnwares and enable them to continuously evolve """ - + def __init__(self, *args, **kwargs): - super(AnchoredMarket, self).__init__(*args, **kwargs) + super(AnchoredOrganizer, self).__init__(*args, **kwargs) self.anchor_learnware_list = {} # anchor_id: anchor learnware def _update_anchor_learnware(self, anchor_id: str, anchor_learnware: Learnware): @@ -101,27 +87,11 @@ class AnchoredMarket(LearnwareMarket): """ pass - def search_anchor_learnware(self, user_info: AnchoredUserInfo) -> Tuple[Any, List[Learnware]]: - """Search anchor Learnwares from anchor_learnware_list based on user_info - - Parameters - ---------- - user_info : AnchoredUserInfo - - user_info with semantic specifications and statistical information - - some statistical information calculated on previous anchor learnwares - - Returns - ------- - Tuple[Any, List[Learnware]]: - return two items: - - - first is the usage of anchor learnwares, e.g., how to use anchors to calculate some statistical information - - second is a list of anchor learnwares - """ - pass - def search_learnware(self, user_info: AnchoredUserInfo) -> Tuple[Any, List[Learnware]]: - """Find helpful learnwares from learnware_list based on user_info + def search_learnware(self, user_info: AnchoredUserInfo, anchored: bool = False) -> Tuple[Any, List[Learnware]]: + """Search learnwares with anchor marget + - if 'anchor' == True, search anchor Learnwares from anchor_learnware_list based on user_info + - if 'anchor' == False, find helpful learnwares from learnware_list based on user_info Parameters ---------- diff --git a/learnware/market/base.py b/learnware/market/base.py index 84a52b4..e927075 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -43,9 +43,6 @@ class BaseUserInfo: return self.stat_info.get(name, None) -class BaseSearchResult: - - pass class LearnwareMarket: """Base interface for market, it provide the interface of search/add/detele/update learnwares""" diff --git a/learnware/market/easymarket/__init__.py b/learnware/market/easy2/__init__.py similarity index 100% rename from learnware/market/easymarket/__init__.py rename to learnware/market/easy2/__init__.py diff --git a/learnware/market/easymarket/checker.py b/learnware/market/easy2/checker.py similarity index 100% rename from learnware/market/easymarket/checker.py rename to learnware/market/easy2/checker.py diff --git a/learnware/market/easymarket/database_ops.py b/learnware/market/easy2/database_ops.py similarity index 100% rename from learnware/market/easymarket/database_ops.py rename to learnware/market/easy2/database_ops.py diff --git a/learnware/market/easymarket/organizer.py b/learnware/market/easy2/organizer.py similarity index 100% rename from learnware/market/easymarket/organizer.py rename to learnware/market/easy2/organizer.py diff --git a/learnware/market/easymarket/searcher.py b/learnware/market/easy2/searcher.py similarity index 100% rename from learnware/market/easymarket/searcher.py rename to learnware/market/easy2/searcher.py diff --git a/learnware/market/evolve/__init__.py b/learnware/market/evolve/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/learnware/market/evolve.py b/learnware/market/evolve/organizer.py similarity index 65% rename from learnware/market/evolve.py rename to learnware/market/evolve/organizer.py index e9e5cc3..ba54c83 100644 --- a/learnware/market/evolve.py +++ b/learnware/market/evolve/organizer.py @@ -1,21 +1,19 @@ -from typing import Tuple, Any, List, Union, Dict +from typing import List -from .base import LearnwareMarket -from ..learnware import Learnware -from ..specification import BaseStatSpecification +from ..easy2.organizer import EasyOrganizer +from ...learnware import Learnware +from ...specification import BaseStatSpecification +from ...logger import get_module_logger +logger = get_module_logger("evolve_organizer") -class EvolvedMarket(LearnwareMarket): - """Organize learnwares and enable them to continuously evolve - Parameters - ---------- - LearnwareMarket : _type_ - Basic market version +class EvolvedOrganizer(EasyOrganizer): + """Organize learnwares and enable them to continuously evolve """ - + def __init__(self, *args, **kwargs): - super(EvolvedMarket, self).__init__(*args, **kwargs) + super(EvolvedOrganizer, self).__init__(*args, **kwargs) def generate_new_stat_specification(self, learnware: Learnware) -> BaseStatSpecification: """Generate new statistical specification for learnwares @@ -39,4 +37,4 @@ class EvolvedMarket(LearnwareMarket): id_list : List[str] Id list for learnwares """ - pass + pass \ No newline at end of file diff --git a/learnware/market/evolve_anchor/__init__.py b/learnware/market/evolve_anchor/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/learnware/market/evolve_anchor.py b/learnware/market/evolve_anchor/organizer.py similarity index 56% rename from learnware/market/evolve_anchor.py rename to learnware/market/evolve_anchor/organizer.py index 55c8192..b5ad551 100644 --- a/learnware/market/evolve_anchor.py +++ b/learnware/market/evolve_anchor/organizer.py @@ -1,22 +1,18 @@ -from typing import Tuple, Any, List, Union, Dict +from typing import List -from .anchor import AnchoredUserInfo, AnchoredMarket -from .evolve import EvolvedMarket +from ..evolve.organizer import EvolvedOrganizer +from ..anchor.organizer import AnchoredOrganizer, AnchoredUserInfo +from ...logger import get_module_logger +logger = get_module_logger("evolve_anchor_organizer") -class EvolvedAnchoredMarket(AnchoredMarket, EvolvedMarket): - """Organize learnwares with anchors and enable them to continuously evolve - Parameters - ---------- - AnchoredMarket : _type_ - Market version with anchors - EvolvedMarket : _type_ - Market version with evolved learnwares +class EvolveAnchoredOrganizer(AnchoredOrganizer, EvolvedOrganizer): + """Organize learnwares and enable them to continuously evolve """ - + def __init__(self, *args, **kwargs): - super(EvolvedAnchoredMarket, self).__init__(*args, **kwargs) + AnchoredOrganizer.__init__(self, *args, **kwargs) def evolve_anchor_learnware_list(self, anchor_id_list: List[str]): """Enable anchor learnwares to evolve, e.g., new stat_spec diff --git a/learnware/market/hetergeneous/__init__.py b/learnware/market/hetergeneous/__init__.py new file mode 100644 index 0000000..e69de29 diff --git a/learnware/market/heterogeneous_feature.py b/learnware/market/hetergeneous/organizer.py similarity index 85% rename from learnware/market/heterogeneous_feature.py rename to learnware/market/hetergeneous/organizer.py index 8a611aa..86ec2d0 100644 --- a/learnware/market/heterogeneous_feature.py +++ b/learnware/market/hetergeneous/organizer.py @@ -1,8 +1,8 @@ import numpy as np -from typing import Tuple, Any, List, Union, Dict +from typing import List -from .evolve import EvolvedMarket -from ..learnware import Learnware +from ..evolve.organizer import EvolvedOrganizer +from ...learnware import Learnware class MappingFunction: @@ -25,17 +25,12 @@ class MappingFunction: pass -class HeterogeneousFeatureMarket(EvolvedMarket): - """Organize learnwares with heterogeneous feature spaces - - Parameters - ---------- - EvolvedMarket : _type_ - Market version with evolved learnwares +class HeterogeneousOrganizer(EvolvedOrganizer): + """Organize learnwares with heterogeneous feature spaces, organizer version with evolved learnwares """ def __init__(self, *args, **kwargs): - super(HeterogeneousFeatureMarket, self).__init__(*args, **kwargs) + super(HeterogeneousOrganizer, self).__init__(*args, **kwargs) self.mapping_function_list = {} def _mapping_function_list_initialization(self, learnware_list: List[Learnware]): From d406ad29f530b73e4c204cb7ed7dd5cf9bec935e Mon Sep 17 00:00:00 2001 From: bxdd Date: Fri, 27 Oct 2023 03:08:11 +0800 Subject: [PATCH 13/35] black format --- .../pfs/pfs_cross_transfer.py | 4 +- learnware/market/anchor/__init__.py | 2 - learnware/market/anchor/organizer.py | 7 +- learnware/market/base.py | 67 +++++++++-------- learnware/market/easy.py | 29 +++++--- learnware/market/easy2/__init__.py | 4 +- learnware/market/easy2/checker.py | 2 +- learnware/market/easy2/organizer.py | 71 ++++++++++--------- learnware/market/easy2/searcher.py | 40 ++++++----- learnware/market/evolve/organizer.py | 7 +- learnware/market/evolve_anchor/organizer.py | 5 +- learnware/market/hetergeneous/organizer.py | 3 +- 12 files changed, 127 insertions(+), 114 deletions(-) diff --git a/examples/dataset_pfs_workflow/pfs/pfs_cross_transfer.py b/examples/dataset_pfs_workflow/pfs/pfs_cross_transfer.py index 5f69127..93a3fa3 100644 --- a/examples/dataset_pfs_workflow/pfs/pfs_cross_transfer.py +++ b/examples/dataset_pfs_workflow/pfs/pfs_cross_transfer.py @@ -85,9 +85,7 @@ def get_split_errs(algo): split = train_xs.shape[0] - proportion_list[tmp] model.fit( - train_xs[ - split:, - ], + train_xs[split:,], train_ys[split:], eval_set=[(val_xs, val_ys)], early_stopping_rounds=50, diff --git a/learnware/market/anchor/__init__.py b/learnware/market/anchor/__init__.py index 80680a9..39d77a1 100644 --- a/learnware/market/anchor/__init__.py +++ b/learnware/market/anchor/__init__.py @@ -1,5 +1,3 @@ - - class AnchoredUserInfo(BaseUserInfo): """ User Information for searching learnware (add the anchor design) diff --git a/learnware/market/anchor/organizer.py b/learnware/market/anchor/organizer.py index e1213fc..f903f1f 100644 --- a/learnware/market/anchor/organizer.py +++ b/learnware/market/anchor/organizer.py @@ -8,6 +8,7 @@ from ...specification import BaseStatSpecification logger = get_module_logger("evolve_organizer") + class AnchoredUserInfo(BaseUserInfo): """ User Information for searching learnware (add the anchor design) @@ -34,9 +35,8 @@ class AnchoredUserInfo(BaseUserInfo): class AnchoredOrganizer(EasyOrganizer): - """Organize learnwares and enable them to continuously evolve - """ - + """Organize learnwares and enable them to continuously evolve""" + def __init__(self, *args, **kwargs): super(AnchoredOrganizer, self).__init__(*args, **kwargs) self.anchor_learnware_list = {} # anchor_id: anchor learnware @@ -87,7 +87,6 @@ class AnchoredOrganizer(EasyOrganizer): """ pass - def search_learnware(self, user_info: AnchoredUserInfo, anchored: bool = False) -> Tuple[Any, List[Learnware]]: """Search learnwares with anchor marget - if 'anchor' == True, search anchor Learnwares from anchor_learnware_list based on user_info diff --git a/learnware/market/base.py b/learnware/market/base.py index e927075..2e39312 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -10,6 +10,7 @@ from ..logger import get_module_logger logger = get_module_logger("market_base", "INFO") + class BaseUserInfo: """User Information for searching learnware""" @@ -43,16 +44,15 @@ class BaseUserInfo: return self.stat_info.get(name, None) - class LearnwareMarket: """Base interface for market, it provide the interface of search/add/detele/update learnwares""" def __init__( self, market_id: str = None, - organizer: 'LearnwareOrganizer' = None, - checker: 'LearnwareChecker' = None, - searcher: 'LearnwareSearcher' = None, + organizer: "LearnwareOrganizer" = None, + checker: "LearnwareChecker" = None, + searcher: "LearnwareSearcher" = None, ): self.market_id = market_id self.learnware_organizer = LearnwareOrganizer() if organizer is None else organizer @@ -64,7 +64,7 @@ class LearnwareMarket: def reload_market(self, **kwargs) -> bool: self.learnware_organizer.reload_market(**kwargs) - + def check_learnware(self, learnware: Learnware, **kwargs) -> bool: return self.learnware_checker(learnware, **kwargs) @@ -80,30 +80,30 @@ class LearnwareMarket: def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, **kwargs) -> bool: return self.learnware_organizer.update_learnware(id, zip_path=zip_path, semantic_spec=semantic_spec, **kwargs) - def get_learnware_ids(self, top:int = None, **kwargs): + def get_learnware_ids(self, top: int = None, **kwargs): return self.learnware_organizer.get_learnware_ids(top, **kwargs) - - - def get_learnwares(self, top:int = None, **kwargs): + + def get_learnwares(self, top: int = None, **kwargs): return self.learnware_organizer.get_learnwares(top, **kwargs) - + def get_learnware_path_by_ids(self, ids: Union[str, List[str]], **kwargs) -> Union[Learnware, List[Learnware]]: raise self.learnware_organizer.get_learnware_path_by_ids(ids, **kwargs) def get_learnware_by_ids(self, id: Union[str, List[str]], **kwargs) -> Union[Learnware, List[Learnware]]: return self.learnware_organizer.get_learnware_by_ids(id, **kwargs) - + + class LearnwareOrganizer: - def __init__(self, market_id, checker: 'LearnwareChecker' = None): + def __init__(self, market_id, checker: "LearnwareChecker" = None): self.reset(market_id=market_id, checker=checker) - - def reset(self, market_id, checker: 'LearnwareChecker', **kwargs): + + def reset(self, market_id, checker: "LearnwareChecker", **kwargs): self.market_id = market_id self.organizer = checker - + def reload_market(self) -> bool: """Reload the learnware organizer when server restared. - + Returns ------- bool @@ -111,7 +111,7 @@ class LearnwareOrganizer: """ raise NotImplementedError("reload market is Not Implemented") - + def add_learnware(self, zip_path: str, semantic_spec: dict) -> Tuple[str, bool]: """Add a learnware into the market. @@ -141,8 +141,7 @@ class LearnwareOrganizer: """ raise NotImplementedError("add learnware is Not Implemented") - - + def delete_learnware(self, id: str) -> bool: """Delete a learnware from market @@ -173,7 +172,7 @@ class LearnwareOrganizer: id of target learnware. """ raise NotImplementedError("update learnware is Not Implemented") - + def get_learnware_by_ids(self, id: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """ Get Learnware from market by id @@ -192,7 +191,7 @@ class LearnwareOrganizer: - 'None' indicating the target id not found. """ raise NotImplementedError("get_learnware_by_ids is not implemented") - + def get_learnware_path_by_ids(self, ids: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """Get Zipped Learnware file by id @@ -210,8 +209,8 @@ class LearnwareOrganizer: None for Learnware NOT Found. """ raise NotImplementedError("get_learnware_path_by_ids is not implemented") - - def get_learnware_ids(self, top:int = None) -> List[str]: + + def get_learnware_ids(self, top: int = None) -> List[str]: """get the list of learnware ids Parameters @@ -225,9 +224,8 @@ class LearnwareOrganizer: the first top ids """ raise NotImplementedError("get_learnware_ids is not implemented") - - - def get_learnwares(self, top:int = None) -> List[Learnware]: + + def get_learnwares(self, top: int = None) -> List[Learnware]: """get the list of learnwares Parameters @@ -242,13 +240,14 @@ class LearnwareOrganizer: """ raise NotImplementedError("get_learnwares is not implemented") + class LearnwareSearcher: def __init__(self, organizer: LearnwareOrganizer = None): self.learnware_oganizer = organizer - + def reset(self, organizer): self.learnware_oganizer = organizer - + def __call__(self, user_info: BaseUserInfo): """Search learnwares based on user_info @@ -258,19 +257,19 @@ class LearnwareSearcher: user_info contains semantic_spec and stat_info """ raise NotImplementedError("'__call__' method is not implemented in LearnwareSearcher") - + class LearnwareChecker: INVALID_LEARNWARE = -1 NONUSABLE_LEARNWARE = 0 USABLE_LEARWARE = 1 - + def __init__(self, organizer: LearnwareOrganizer = None): self.learnware_oganizer = organizer - + def reset(self, organizer): self.learnware_oganizer = organizer - + def __call__(self, learnware: Learnware) -> int: """Check the utility of a learnware @@ -286,5 +285,5 @@ class LearnwareChecker: - The NOPREDICTION_LEARNWARE denotes the learnware pass the check but cannot make prediction due to some env dependency - The NOPREDICTION_LEARNWARE denotes the leanrware pass the check and can make prediction """ - - raise NotImplementedError("'__call__' method is not implemented in LearnwareChecker") \ No newline at end of file + + raise NotImplementedError("'__call__' method is not implemented in LearnwareChecker") diff --git a/learnware/market/easy.py b/learnware/market/easy.py index 63a357e..591cf05 100644 --- a/learnware/market/easy.py +++ b/learnware/market/easy.py @@ -699,6 +699,7 @@ class EasyMarket(LearnwareMarket): List[Learnware] The list of returned learnwares """ + def _match_semantic_spec_tag(semantic_spec1, semantic_spec2) -> bool: """Judge if tags of two semantic specs are consistent @@ -737,8 +738,8 @@ class EasyMarket(LearnwareMarket): elif semantic_spec1[key]["Type"] == "Tag": if not (set(v1) & set(v2)): return False - return True - + return True + matched_learnware_tag = [] final_result = [] user_semantic_spec = user_info.get_semantic_spec() @@ -747,15 +748,21 @@ class EasyMarket(LearnwareMarket): learnware_semantic_spec = learnware.get_specification().get_semantic_spec() if _match_semantic_spec_tag(user_semantic_spec, learnware_semantic_spec): matched_learnware_tag.append(learnware) - + if len(matched_learnware_tag) > 0: if "Name" in user_semantic_spec: name_user = user_semantic_spec["Name"]["Values"].lower() if len(name_user) > 0: # Exact search - name_list = [learnware.get_specification().get_semantic_spec()["Name"]["Values"].lower() for learnware in matched_learnware_tag] - des_list = [learnware.get_specification().get_semantic_spec()["Description"]["Values"].lower() for learnware in matched_learnware_tag] - + name_list = [ + learnware.get_specification().get_semantic_spec()["Name"]["Values"].lower() + for learnware in matched_learnware_tag + ] + des_list = [ + learnware.get_specification().get_semantic_spec()["Description"]["Values"].lower() + for learnware in matched_learnware_tag + ] + matched_learnware_exact = [] for i in range(len(name_list)): if name_user in name_list[i] or name_user in des_list[i]: @@ -771,9 +778,11 @@ class EasyMarket(LearnwareMarket): if final_score >= min_score: matched_learnware_fuzz.append(matched_learnware_tag[i]) fuzz_scores.append(final_score) - + # Sort by score - sort_idx = sorted(list(range(len(fuzz_scores))), key=lambda k: fuzz_scores[k], reverse=True)[:max_num] + sort_idx = sorted(list(range(len(fuzz_scores))), key=lambda k: fuzz_scores[k], reverse=True)[ + :max_num + ] final_result = [matched_learnware_fuzz[idx] for idx in sort_idx] else: final_result = matched_learnware_exact @@ -782,9 +791,7 @@ class EasyMarket(LearnwareMarket): else: final_result = matched_learnware_tag - logger.info( - "semantic_spec search: choose %d from %d learnwares" % (len(final_result), len(learnware_list)) - ) + logger.info("semantic_spec search: choose %d from %d learnwares" % (len(final_result), len(learnware_list))) return final_result def search_learnware( diff --git a/learnware/market/easy2/__init__.py b/learnware/market/easy2/__init__.py index 96f0a34..722e93d 100644 --- a/learnware/market/easy2/__init__.py +++ b/learnware/market/easy2/__init__.py @@ -1,7 +1,9 @@ from ..base import LearnwareSearcher, LearnwareOrganizer + class EasySearcher(LearnwareSearcher): pass + class EasyOrganizer(LearnwareOrganizer): - pass \ No newline at end of file + pass diff --git a/learnware/market/easy2/checker.py b/learnware/market/easy2/checker.py index ce5d2b4..8a4a250 100644 --- a/learnware/market/easy2/checker.py +++ b/learnware/market/easy2/checker.py @@ -7,8 +7,8 @@ from ...logger import get_module_logger logger = get_module_logger("easy_checker", "INFO") + class EasyChecker(LearnwareChecker): - def __call__(self, learnware): semantic_spec = learnware.get_specification().get_semantic_spec() diff --git a/learnware/market/easy2/organizer.py b/learnware/market/easy2/organizer.py index 2295f63..c5d9629 100644 --- a/learnware/market/easy2/organizer.py +++ b/learnware/market/easy2/organizer.py @@ -30,23 +30,22 @@ logger = get_module_logger("easy_organizer") class EasyOrganizer(LearnwareOrganizer): - - def __init__(self, market_id, checker: 'EasyChecker' = None, rebuild: bool = False): + def __init__(self, market_id, checker: "EasyChecker" = None, rebuild: bool = False): self.reset(market_id=market_id, checker=checker, rebuild=rebuild) def reset(self, market_id, checker: EasyChecker = None, rebuild: bool = False): super(EasyOrganizer, self).reset(market_id=market_id, checker=checker) self.reload_market(rebuild=rebuild) - + def reload_market(self, rebuild=False) -> bool: """Reload the learnware organizer when server restared. - + Returns ------- bool A flag indicating whether the market is reload successfully. """ - + self.market_store_path = os.path.join(conf.market_root_path, self.market_id) self.learnware_pool_path = os.path.join(self.market_store_path, "learnware_pool") self.learnware_zip_pool_path = os.path.join(self.learnware_pool_path, "zips") @@ -58,7 +57,7 @@ class EasyOrganizer(LearnwareOrganizer): self.count = 0 self.semantic_spec_list = conf.semantic_specs self.dbops = DatabaseOperations(conf.database_url, "market_" + self.market_id) - + if rebuild: logger.warning("Warning! You are trying to clear current database!") try: @@ -70,10 +69,17 @@ class EasyOrganizer(LearnwareOrganizer): os.makedirs(self.learnware_pool_path, exist_ok=True) os.makedirs(self.learnware_zip_pool_path, exist_ok=True) os.makedirs(self.learnware_folder_pool_path, exist_ok=True) - self.learnware_list, self.learnware_zip_list, self.learnware_folder_list, self.use_flags, self.count = self.dbops.load_market() - - - def add_learnware(self, zip_path: str, semantic_spec: dict, id: str = None, check_status: int = None) -> Tuple[str, bool]: + ( + self.learnware_list, + self.learnware_zip_list, + self.learnware_folder_list, + self.use_flags, + self.count, + ) = self.dbops.load_market() + + def add_learnware( + self, zip_path: str, semantic_spec: dict, id: str = None, check_status: int = None + ) -> Tuple[str, bool]: """Add a learnware into the market. .. note:: @@ -133,7 +139,7 @@ class EasyOrganizer(LearnwareOrganizer): return None, EasyChecker.INVALID_LEARNWARE logger.info("Get new learnware from %s" % (zip_path)) - + id = id if id is not None else "%08d" % (self.count) target_zip_dir = os.path.join(self.learnware_zip_pool_path, "%s.zip" % (id)) target_folder_dir = os.path.join(self.learnware_folder_pool_path, id) @@ -159,7 +165,7 @@ class EasyOrganizer(LearnwareOrganizer): return None, EasyChecker.INVALID_LEARNWARE learnwere_status = check_status if check_status is not None else self.checker.check_learnware(new_learnware) - + self.dbops.add_learnware( id=id, semantic_spec=semantic_spec, @@ -205,9 +211,7 @@ class EasyOrganizer(LearnwareOrganizer): return True - def update_learnware( - self, id: str, zip_path: str = None, semantic_spec: dict = None, check_status: int = None - ): + def update_learnware(self, id: str, zip_path: str = None, semantic_spec: dict = None, check_status: int = None): """update learnware with zip_path and semantic_specification TODO: update should pass the semantic check too @@ -227,17 +231,21 @@ class EasyOrganizer(LearnwareOrganizer): _type_ _description_ """ - assert zip_path is None and semantic_spec is None, f"at least one of 'zip_path' and 'semantic_spec' should not be None when update learnware" + assert ( + zip_path is None and semantic_spec is None + ), f"at least one of 'zip_path' and 'semantic_spec' should not be None when update learnware" assert check_status != EasyChecker.INVALID_LEARNWARE, f"'check_status' can not be INVALID_LEARNWARE" if zip_path is None and check_status is not None: logger.warning("check_status will be ignored when zip_path is None for learnware update") - + learnware_zippath = self.learnware_zip_list[id] if zip_path is None else zip_path - semantic_spec = self.learnware_list[id].get_specification().get_semantic_spec() if semantic_spec is None else semantic_spec - + semantic_spec = ( + self.learnware_list[id].get_specification().get_semantic_spec() if semantic_spec is None else semantic_spec + ) + self.dbops.update_learnware_semantic_specification(id, semantic_spec) - + target_zip_dir = self.learnware_zip_list[id] target_folder_dir = self.learnware_folder_list[id] @@ -245,25 +253,25 @@ class EasyOrganizer(LearnwareOrganizer): with tempfile.TemporaryDirectory(prefix="learnware_") as tempdir: with zipfile.ZipFile(zip_path, "r") as z_file: z_file.extractall(tempdir) - + try: new_learnware = get_learnware_from_dirpath( id=id, semantic_spec=semantic_spec, learnware_dirpath=tempdir ) except Exception: return EasyChecker.INVALID_LEARNWARE - + if new_learnware is None: return EasyChecker.INVALID_LEARNWARE - + learnwere_status = self.checker.check_learnware(new_learnware) else: learnwere_status = self.use_flags[id] if zip_path is None else check_status - + copyfile(zip_path, target_zip_dir) with zipfile.ZipFile(target_zip_dir, "r") as z_file: - z_file.extractall(target_folder_dir) - + z_file.extractall(target_folder_dir) + self.learnware_list[id] = get_learnware_from_dirpath( id=id, semantic_spec=semantic_spec, learnware_dirpath=target_folder_dir ) @@ -271,7 +279,6 @@ class EasyOrganizer(LearnwareOrganizer): self.dbops.update_learnware_use_flag(id, learnwere_status) return learnwere_status - def get_learnware_by_ids(self, ids: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """Search learnware by id or list of ids. @@ -335,17 +342,15 @@ class EasyOrganizer(LearnwareOrganizer): except: logger.warning("Learnware ID '%s' NOT Found!" % (ids)) return None - - def get_learnware_ids(self, top:int = None) -> List[str]: + + def get_learnware_ids(self, top: int = None) -> List[str]: if top is None: return list(self.learnware_list.keys()) else: return list(self.learnware_list.keys())[:top] - - - def get_learnwares(self, top:int = None) -> List[str]: + + def get_learnwares(self, top: int = None) -> List[str]: if top is None: return list(self.learnware_list.values()) else: return list(self.learnware_list.values())[:top] - \ No newline at end of file diff --git a/learnware/market/easy2/searcher.py b/learnware/market/easy2/searcher.py index e5adcd8..2a64f9d 100644 --- a/learnware/market/easy2/searcher.py +++ b/learnware/market/easy2/searcher.py @@ -9,10 +9,10 @@ from ...learnware import Learnware from ...specification import RKMEStatSpecification from ...logger import get_module_logger -logger = get_module_logger('easy_seacher') +logger = get_module_logger("easy_seacher") + class EasySearcher(LearnwareSearcher): - def _convert_dist_to_score( self, dist_list: List[float], dist_epsilon: float = 0.01, min_score: float = 0.92 ) -> List[float]: @@ -454,6 +454,7 @@ class EasySearcher(LearnwareSearcher): List[Learnware] The list of returned learnwares """ + def _match_semantic_spec_tag(semantic_spec1, semantic_spec2) -> bool: """Judge if tags of two semantic specs are consistent @@ -492,8 +493,8 @@ class EasySearcher(LearnwareSearcher): elif semantic_spec1[key]["Type"] == "Tag": if not (set(v1) & set(v2)): return False - return True - + return True + matched_learnware_tag = [] final_result = [] user_semantic_spec = user_info.get_semantic_spec() @@ -502,15 +503,21 @@ class EasySearcher(LearnwareSearcher): learnware_semantic_spec = learnware.get_specification().get_semantic_spec() if _match_semantic_spec_tag(user_semantic_spec, learnware_semantic_spec): matched_learnware_tag.append(learnware) - + if len(matched_learnware_tag) > 0: if "Name" in user_semantic_spec: name_user = user_semantic_spec["Name"]["Values"].lower() if len(name_user) > 0: # Exact search - name_list = [learnware.get_specification().get_semantic_spec()["Name"]["Values"].lower() for learnware in matched_learnware_tag] - des_list = [learnware.get_specification().get_semantic_spec()["Description"]["Values"].lower() for learnware in matched_learnware_tag] - + name_list = [ + learnware.get_specification().get_semantic_spec()["Name"]["Values"].lower() + for learnware in matched_learnware_tag + ] + des_list = [ + learnware.get_specification().get_semantic_spec()["Description"]["Values"].lower() + for learnware in matched_learnware_tag + ] + matched_learnware_exact = [] for i in range(len(name_list)): if name_user in name_list[i] or name_user in des_list[i]: @@ -526,9 +533,11 @@ class EasySearcher(LearnwareSearcher): if final_score >= min_score: matched_learnware_fuzz.append(matched_learnware_tag[i]) fuzz_scores.append(final_score) - + # Sort by score - sort_idx = sorted(list(range(len(fuzz_scores))), key=lambda k: fuzz_scores[k], reverse=True)[:max_num] + sort_idx = sorted(list(range(len(fuzz_scores))), key=lambda k: fuzz_scores[k], reverse=True)[ + :max_num + ] final_result = [matched_learnware_fuzz[idx] for idx in sort_idx] else: final_result = matched_learnware_exact @@ -537,12 +546,12 @@ class EasySearcher(LearnwareSearcher): else: final_result = matched_learnware_tag - logger.info( - "semantic_spec search: choose %d from %d learnwares" % (len(final_result), len(learnware_list)) - ) + logger.info("semantic_spec search: choose %d from %d learnwares" % (len(final_result), len(learnware_list))) return final_result - - def __call__(self, user_info: BaseUserInfo, max_search_num: int = 5, search_method: str = "greedy") -> Tuple[List[float], List[Learnware], float, List[Learnware]]: + + def __call__( + self, user_info: BaseUserInfo, max_search_num: int = 5, search_method: str = "greedy" + ) -> Tuple[List[float], List[Learnware], float, List[Learnware]]: """Search learnwares based on user_info Parameters @@ -605,4 +614,3 @@ class EasySearcher(LearnwareSearcher): logger.info(f"After filter by rkme spec, learnware_list length is {len(learnware_list)}") return sorted_score_list, single_learnware_list, mixture_score, mixture_learnware_list - \ No newline at end of file diff --git a/learnware/market/evolve/organizer.py b/learnware/market/evolve/organizer.py index ba54c83..7976542 100644 --- a/learnware/market/evolve/organizer.py +++ b/learnware/market/evolve/organizer.py @@ -9,9 +9,8 @@ logger = get_module_logger("evolve_organizer") class EvolvedOrganizer(EasyOrganizer): - """Organize learnwares and enable them to continuously evolve - """ - + """Organize learnwares and enable them to continuously evolve""" + def __init__(self, *args, **kwargs): super(EvolvedOrganizer, self).__init__(*args, **kwargs) @@ -37,4 +36,4 @@ class EvolvedOrganizer(EasyOrganizer): id_list : List[str] Id list for learnwares """ - pass \ No newline at end of file + pass diff --git a/learnware/market/evolve_anchor/organizer.py b/learnware/market/evolve_anchor/organizer.py index b5ad551..1e8173e 100644 --- a/learnware/market/evolve_anchor/organizer.py +++ b/learnware/market/evolve_anchor/organizer.py @@ -8,9 +8,8 @@ logger = get_module_logger("evolve_anchor_organizer") class EvolveAnchoredOrganizer(AnchoredOrganizer, EvolvedOrganizer): - """Organize learnwares and enable them to continuously evolve - """ - + """Organize learnwares and enable them to continuously evolve""" + def __init__(self, *args, **kwargs): AnchoredOrganizer.__init__(self, *args, **kwargs) diff --git a/learnware/market/hetergeneous/organizer.py b/learnware/market/hetergeneous/organizer.py index 86ec2d0..90e2410 100644 --- a/learnware/market/hetergeneous/organizer.py +++ b/learnware/market/hetergeneous/organizer.py @@ -26,8 +26,7 @@ class MappingFunction: class HeterogeneousOrganizer(EvolvedOrganizer): - """Organize learnwares with heterogeneous feature spaces, organizer version with evolved learnwares - """ + """Organize learnwares with heterogeneous feature spaces, organizer version with evolved learnwares""" def __init__(self, *args, **kwargs): super(HeterogeneousOrganizer, self).__init__(*args, **kwargs) From 4b50737cd7210704387055903c0f6a459ae8bdda Mon Sep 17 00:00:00 2001 From: bxdd Date: Fri, 27 Oct 2023 03:19:41 +0800 Subject: [PATCH 14/35] [MNT] add market module --- learnware/market/__init__.py | 14 +++++---- learnware/market/anchor/__init__.py | 36 +--------------------- learnware/market/easy2/__init__.py | 12 ++------ learnware/market/evolve/__init__.py | 1 + learnware/market/evolve_anchor/__init__.py | 1 + learnware/market/hetergeneous/__init__.py | 1 + learnware/market/module.py | 19 ++++++++++++ 7 files changed, 34 insertions(+), 50 deletions(-) diff --git a/learnware/market/__init__.py b/learnware/market/__init__.py index bf18990..ff30c4b 100644 --- a/learnware/market/__init__.py +++ b/learnware/market/__init__.py @@ -1,6 +1,8 @@ -from .anchor import AnchoredUserInfo, AnchoredMarket -from .base import BaseUserInfo, LearnwareMarket -from .evolve_anchor import EvolvedAnchoredMarket -from .evolve import EvolvedMarket -from .easy import EasyMarket -from .heterogeneous_feature import HeterogeneousFeatureMarket +from .anchor import AnchoredUserInfo, AnchoredOrganizer +from .base import BaseUserInfo, LearnwareMarket, LearnwareChecker, LearnwareOrganizer +from .evolve_anchor import EvolveAnchoredOrganizer +from .evolve import EvolvedOrganizer +from .easy2 import EasyChecker, EasyOrganizer, EasySearcher +from .hetergeneous import HeterogeneousOrganizer, MappingFunction + +from .module import instatiate_learnware_market \ No newline at end of file diff --git a/learnware/market/anchor/__init__.py b/learnware/market/anchor/__init__.py index 39d77a1..66b07d2 100644 --- a/learnware/market/anchor/__init__.py +++ b/learnware/market/anchor/__init__.py @@ -1,35 +1 @@ -class AnchoredUserInfo(BaseUserInfo): - """ - User Information for searching learnware (add the anchor design) - - - UserInfo contains the anchor list acquired from the market - - UserInfo can update stat_info based on anchors - """ - - def __init__(self, id: str, semantic_spec: dict = dict(), stat_info: dict = dict()): - super(AnchoredUserInfo, self).__init__(id, semantic_spec, stat_info) - self.anchor_learnware_list = {} # id: Learnware - - def add_anchor_learnware(self, learnware_id: str, learnware: Learnware): - """Add the anchor learnware acquired from the market - - Parameters - ---------- - learnware_id : str - Id of anchor learnware - learnware : Learnware - Anchor learnware for capturing user requirements - """ - self.anchor_learnware_list[learnware_id] = learnware - - def update_stat_info(self, name: str, item: Any): - """Update stat_info based on anchor learnwares - - Parameters - ---------- - name : str - Name of stat_info - item : Any - Statistical information calculated on anchor learnwares - """ - self.stat_info[name] = item +from .organizer import AnchoredOrganizer, AnchoredUserInfo \ No newline at end of file diff --git a/learnware/market/easy2/__init__.py b/learnware/market/easy2/__init__.py index 722e93d..ef05bb9 100644 --- a/learnware/market/easy2/__init__.py +++ b/learnware/market/easy2/__init__.py @@ -1,9 +1,3 @@ -from ..base import LearnwareSearcher, LearnwareOrganizer - - -class EasySearcher(LearnwareSearcher): - pass - - -class EasyOrganizer(LearnwareOrganizer): - pass +from .organizer import EasyOrganizer +from .checker import EasyChecker +from .searcher import EasySearcher \ No newline at end of file diff --git a/learnware/market/evolve/__init__.py b/learnware/market/evolve/__init__.py index e69de29..c096db0 100644 --- a/learnware/market/evolve/__init__.py +++ b/learnware/market/evolve/__init__.py @@ -0,0 +1 @@ +from .organizer import EvolvedOrganizer \ No newline at end of file diff --git a/learnware/market/evolve_anchor/__init__.py b/learnware/market/evolve_anchor/__init__.py index e69de29..030b444 100644 --- a/learnware/market/evolve_anchor/__init__.py +++ b/learnware/market/evolve_anchor/__init__.py @@ -0,0 +1 @@ +from .organizer import EvolveAnchoredOrganizer \ No newline at end of file diff --git a/learnware/market/hetergeneous/__init__.py b/learnware/market/hetergeneous/__init__.py index e69de29..77943cd 100644 --- a/learnware/market/hetergeneous/__init__.py +++ b/learnware/market/hetergeneous/__init__.py @@ -0,0 +1 @@ +from .organizer import MappingFunction, HeterogeneousOrganizer \ No newline at end of file diff --git a/learnware/market/module.py b/learnware/market/module.py index e69de29..fbd5978 100644 --- a/learnware/market/module.py +++ b/learnware/market/module.py @@ -0,0 +1,19 @@ +from .base import LearnwareMarket +from .easy2 import EasyChecker, EasyOrganizer, EasySearcher + +MARKET_CONFIG = { + 'easy': { + 'organizer': EasyOrganizer(), + 'checker': EasyChecker(), + 'searcher': EasySearcher(), + } +} + +def instatiate_learnware_market(market_id, name='easy'): + + return LearnwareMarket( + market_id=market_id, + organizer=MARKET_CONFIG[name]['organizer'], + checker=MARKET_CONFIG[name]['checker'], + searcher=MARKET_CONFIG[name]['searcher'], + ) \ No newline at end of file From bef000fb0a14d3f3ea1ddd18043c22835e469c03 Mon Sep 17 00:00:00 2001 From: bxdd Date: Fri, 27 Oct 2023 03:20:00 +0800 Subject: [PATCH 15/35] [MNT] black format --- learnware/market/__init__.py | 2 +- learnware/market/anchor/__init__.py | 2 +- learnware/market/easy2/__init__.py | 2 +- learnware/market/evolve/__init__.py | 2 +- learnware/market/evolve_anchor/__init__.py | 2 +- learnware/market/hetergeneous/__init__.py | 2 +- learnware/market/module.py | 20 ++++++++++---------- 7 files changed, 16 insertions(+), 16 deletions(-) diff --git a/learnware/market/__init__.py b/learnware/market/__init__.py index ff30c4b..9201de7 100644 --- a/learnware/market/__init__.py +++ b/learnware/market/__init__.py @@ -5,4 +5,4 @@ from .evolve import EvolvedOrganizer from .easy2 import EasyChecker, EasyOrganizer, EasySearcher from .hetergeneous import HeterogeneousOrganizer, MappingFunction -from .module import instatiate_learnware_market \ No newline at end of file +from .module import instatiate_learnware_market diff --git a/learnware/market/anchor/__init__.py b/learnware/market/anchor/__init__.py index 66b07d2..64e0622 100644 --- a/learnware/market/anchor/__init__.py +++ b/learnware/market/anchor/__init__.py @@ -1 +1 @@ -from .organizer import AnchoredOrganizer, AnchoredUserInfo \ No newline at end of file +from .organizer import AnchoredOrganizer, AnchoredUserInfo diff --git a/learnware/market/easy2/__init__.py b/learnware/market/easy2/__init__.py index ef05bb9..2ab8c48 100644 --- a/learnware/market/easy2/__init__.py +++ b/learnware/market/easy2/__init__.py @@ -1,3 +1,3 @@ from .organizer import EasyOrganizer from .checker import EasyChecker -from .searcher import EasySearcher \ No newline at end of file +from .searcher import EasySearcher diff --git a/learnware/market/evolve/__init__.py b/learnware/market/evolve/__init__.py index c096db0..e0069c5 100644 --- a/learnware/market/evolve/__init__.py +++ b/learnware/market/evolve/__init__.py @@ -1 +1 @@ -from .organizer import EvolvedOrganizer \ No newline at end of file +from .organizer import EvolvedOrganizer diff --git a/learnware/market/evolve_anchor/__init__.py b/learnware/market/evolve_anchor/__init__.py index 030b444..18624c8 100644 --- a/learnware/market/evolve_anchor/__init__.py +++ b/learnware/market/evolve_anchor/__init__.py @@ -1 +1 @@ -from .organizer import EvolveAnchoredOrganizer \ No newline at end of file +from .organizer import EvolveAnchoredOrganizer diff --git a/learnware/market/hetergeneous/__init__.py b/learnware/market/hetergeneous/__init__.py index 77943cd..caef8fa 100644 --- a/learnware/market/hetergeneous/__init__.py +++ b/learnware/market/hetergeneous/__init__.py @@ -1 +1 @@ -from .organizer import MappingFunction, HeterogeneousOrganizer \ No newline at end of file +from .organizer import MappingFunction, HeterogeneousOrganizer diff --git a/learnware/market/module.py b/learnware/market/module.py index fbd5978..80da5d6 100644 --- a/learnware/market/module.py +++ b/learnware/market/module.py @@ -2,18 +2,18 @@ from .base import LearnwareMarket from .easy2 import EasyChecker, EasyOrganizer, EasySearcher MARKET_CONFIG = { - 'easy': { - 'organizer': EasyOrganizer(), - 'checker': EasyChecker(), - 'searcher': EasySearcher(), + "easy": { + "organizer": EasyOrganizer(), + "checker": EasyChecker(), + "searcher": EasySearcher(), } } -def instatiate_learnware_market(market_id, name='easy'): - + +def instatiate_learnware_market(market_id, name="easy"): return LearnwareMarket( market_id=market_id, - organizer=MARKET_CONFIG[name]['organizer'], - checker=MARKET_CONFIG[name]['checker'], - searcher=MARKET_CONFIG[name]['searcher'], - ) \ No newline at end of file + organizer=MARKET_CONFIG[name]["organizer"], + checker=MARKET_CONFIG[name]["checker"], + searcher=MARKET_CONFIG[name]["searcher"], + ) From f2f5c1c0a14c31e2696ae497131cbbeab0831034 Mon Sep 17 00:00:00 2001 From: bxdd Date: Sat, 28 Oct 2023 14:41:33 +0800 Subject: [PATCH 16/35] [FIX] Fix many bugs --- learnware/market/__init__.py | 1 + learnware/market/base.py | 12 +- learnware/market/easy2/database_ops.py | 4 +- learnware/market/easy2/organizer.py | 8 +- tests/test_market/test_easy.py | 220 +++++++++++++++++++++++++ 5 files changed, 231 insertions(+), 14 deletions(-) create mode 100644 tests/test_market/test_easy.py diff --git a/learnware/market/__init__.py b/learnware/market/__init__.py index 9201de7..81a4184 100644 --- a/learnware/market/__init__.py +++ b/learnware/market/__init__.py @@ -5,4 +5,5 @@ from .evolve import EvolvedOrganizer from .easy2 import EasyChecker, EasyOrganizer, EasySearcher from .hetergeneous import HeterogeneousOrganizer, MappingFunction +from .easy import EasyMarket from .module import instatiate_learnware_market diff --git a/learnware/market/base.py b/learnware/market/base.py index 2e39312..48dfd6b 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -53,12 +53,14 @@ class LearnwareMarket: organizer: "LearnwareOrganizer" = None, checker: "LearnwareChecker" = None, searcher: "LearnwareSearcher" = None, + rebuild=False, ): self.market_id = market_id self.learnware_organizer = LearnwareOrganizer() if organizer is None else organizer self.learnware_checker = LearnwareChecker() if checker is None else checker self.learnware_checker.reset(organizer=self.learnware_organizer) self.learnware_organizer.reset(market_id=market_id, checker=self.learnware_checker) + self.learnware_organizer.reload_market(rebuild=rebuild) self.learnware_searcher = LearnwareSearcher() if searcher is None else searcher self.learnware_searcher.reset(organizer=self.learnware_organizer) @@ -94,14 +96,14 @@ class LearnwareMarket: class LearnwareOrganizer: - def __init__(self, market_id, checker: "LearnwareChecker" = None): + def __init__(self, market_id=None, checker: 'LearnwareChecker' = None): self.reset(market_id=market_id, checker=checker) - - def reset(self, market_id, checker: "LearnwareChecker", **kwargs): + + def reset(self, market_id=None, checker: 'LearnwareChecker'=None, **kwargs): self.market_id = market_id - self.organizer = checker + self.checker = checker - def reload_market(self) -> bool: + def reload_market(self, rebuild=False, **kwargs) -> bool: """Reload the learnware organizer when server restared. Returns diff --git a/learnware/market/easy2/database_ops.py b/learnware/market/easy2/database_ops.py index 61bc02d..25f02e9 100644 --- a/learnware/market/easy2/database_ops.py +++ b/learnware/market/easy2/database_ops.py @@ -3,8 +3,8 @@ from sqlalchemy import create_engine, text from sqlalchemy import Column, Integer, Text, DateTime, String import os import json -from ..learnware import get_learnware_from_dirpath -from ..logger import get_module_logger +from ...learnware import get_learnware_from_dirpath +from ...logger import get_module_logger logger = get_module_logger("database") DeclarativeBase = declarative_base() diff --git a/learnware/market/easy2/organizer.py b/learnware/market/easy2/organizer.py index c5d9629..09a0f3d 100644 --- a/learnware/market/easy2/organizer.py +++ b/learnware/market/easy2/organizer.py @@ -30,13 +30,7 @@ logger = get_module_logger("easy_organizer") class EasyOrganizer(LearnwareOrganizer): - def __init__(self, market_id, checker: "EasyChecker" = None, rebuild: bool = False): - self.reset(market_id=market_id, checker=checker, rebuild=rebuild) - - def reset(self, market_id, checker: EasyChecker = None, rebuild: bool = False): - super(EasyOrganizer, self).reset(market_id=market_id, checker=checker) - self.reload_market(rebuild=rebuild) - + def reload_market(self, rebuild=False) -> bool: """Reload the learnware organizer when server restared. diff --git a/tests/test_market/test_easy.py b/tests/test_market/test_easy.py new file mode 100644 index 0000000..38459ad --- /dev/null +++ b/tests/test_market/test_easy.py @@ -0,0 +1,220 @@ +import sys +import unittest +import os +import copy +import joblib +import zipfile +import numpy as np +from sklearn import svm +from sklearn.datasets import load_digits +from sklearn.model_selection import train_test_split +from shutil import copyfile, rmtree + +import learnware +from learnware.market import EasyMarket, BaseUserInfo +from learnware.learnware import JobSelectorReuser, AveragingReuser, EnsemblePruningReuser +import learnware.specification as specification + +curr_root = os.path.dirname(os.path.abspath(__file__)) + +user_semantic = { + "Data": {"Values": ["Image"], "Type": "Class"}, + "Task": { + "Values": ["Classification"], + "Type": "Class", + }, + "Library": {"Values": ["Scikit-learn"], "Type": "Class"}, + "Scenario": {"Values": ["Education"], "Type": "Tag"}, + "Description": {"Values": "", "Type": "String"}, + "Name": {"Values": "", "Type": "String"}, +} + + +class TestAllWorkflow(unittest.TestCase): + @classmethod + def setUpClass(cls) -> None: + np.random.seed(2023) + learnware.init() + + def _init_learnware_market(self): + """initialize learnware market""" + easy_market = EasyMarket(market_id="sklearn_digits", rebuild=True) + return easy_market + + def test_prepare_learnware_randomly(self, learnware_num=5): + self.zip_path_list = [] + X, y = load_digits(return_X_y=True) + + for i in range(learnware_num): + dir_path = os.path.join(curr_root, "learnware_pool", "svm_%d" % (i)) + os.makedirs(dir_path, exist_ok=True) + + print("Preparing Learnware: %d" % (i)) + + data_X, _, data_y, _ = train_test_split(X, y, test_size=0.3, shuffle=True) + clf = svm.SVC(kernel="linear", probability=True) + clf.fit(data_X, data_y) + + 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.save(os.path.join(dir_path, "svm.json")) + + init_file = os.path.join(dir_path, "__init__.py") + copyfile( + os.path.join(curr_root, "learnware_example/example_init.py"), init_file + ) # cp example_init.py init_file + + yaml_file = os.path.join(dir_path, "learnware.yaml") + copyfile(os.path.join(curr_root, "learnware_example/example.yaml"), yaml_file) # cp example.yaml yaml_file + + env_file = os.path.join(dir_path, "environment.yaml") + copyfile(os.path.join(curr_root, "learnware_example/environment.yaml"), env_file) + + zip_file = dir_path + ".zip" + # zip -q -r -j zip_file dir_path + with zipfile.ZipFile(zip_file, "w") as zip_obj: + for foldername, subfolders, filenames in os.walk(dir_path): + for filename in filenames: + file_path = os.path.join(foldername, filename) + zip_info = zipfile.ZipInfo(filename) + zip_info.compress_type = zipfile.ZIP_STORED + with open(file_path, "rb") as file: + zip_obj.writestr(zip_info, file.read()) + + rmtree(dir_path) # rm -r dir_path + + self.zip_path_list.append(zip_file) + + def test_upload_delete_learnware(self, learnware_num=5, delete=False): + easy_market = self._init_learnware_market() + self.test_prepare_learnware_randomly(learnware_num) + + print("Total Item:", len(easy_market)) + + for idx, zip_path in enumerate(self.zip_path_list): + semantic_spec = copy.deepcopy(user_semantic) + semantic_spec["Name"]["Values"] = "learnware_%d" % (idx) + semantic_spec["Description"]["Values"] = "test_learnware_number_%d" % (idx) + semantic_spec["Output"] = {"Dimension": 1, "Description": {"0": "The label of the hand-written digit."}} + easy_market.add_learnware(zip_path, semantic_spec) + + print("Total Item:", len(easy_market)) + curr_inds = easy_market._get_ids() + print("Available ids After Uploading Learnwares:", curr_inds) + + if delete: + for learnware_id in curr_inds: + easy_market.delete_learnware(learnware_id) + curr_inds = easy_market._get_ids() + print("Available ids After Deleting Learnwares:", curr_inds) + + return easy_market + + def test_search_semantics(self, learnware_num=5): + easy_market = self.test_upload_delete_learnware(learnware_num, delete=False) + print("Total Item:", len(easy_market)) + + test_folder = os.path.join(curr_root, "test_semantics") + + # unzip -o -q zip_path -d unzip_dir + if os.path.exists(test_folder): + rmtree(test_folder) + os.makedirs(test_folder, exist_ok=True) + + with zipfile.ZipFile(self.zip_path_list[0], "r") as zip_obj: + zip_obj.extractall(path=test_folder) + + semantic_spec = copy.deepcopy(user_semantic) + semantic_spec["Name"]["Values"] = f"learnware_{learnware_num - 1}" + semantic_spec["Description"]["Values"] = f"test_learnware_number_{learnware_num - 1}" + + user_info = BaseUserInfo(semantic_spec=semantic_spec) + _, single_learnware_list, _, _ = easy_market.search_learnware(user_info) + + print("User info:", user_info.get_semantic_spec()) + print(f"Search result:") + for learnware in single_learnware_list: + print("Choose learnware:", learnware.id, learnware.get_specification().get_semantic_spec()) + + rmtree(test_folder) # rm -r test_folder + + def test_stat_search(self, learnware_num=5): + easy_market = self.test_upload_delete_learnware(learnware_num, delete=False) + print("Total Item:", len(easy_market)) + + test_folder = os.path.join(curr_root, "test_stat") + + for idx, zip_path in enumerate(self.zip_path_list): + unzip_dir = os.path.join(test_folder, f"{idx}") + + # unzip -o -q zip_path -d unzip_dir + if os.path.exists(unzip_dir): + rmtree(unzip_dir) + os.makedirs(unzip_dir, exist_ok=True) + with zipfile.ZipFile(zip_path, "r") as zip_obj: + zip_obj.extractall(path=unzip_dir) + + user_spec = specification.rkme.RKMEStatSpecification() + user_spec.load(os.path.join(unzip_dir, "svm.json")) + user_info = BaseUserInfo(semantic_spec=user_semantic, stat_info={"RKMEStatSpecification": user_spec}) + ( + sorted_score_list, + single_learnware_list, + mixture_score, + mixture_learnware_list, + ) = easy_market.search_learnware(user_info) + + print(f"search result of user{idx}:") + for score, learnware in zip(sorted_score_list, single_learnware_list): + print(f"score: {score}, learnware_id: {learnware.id}") + print(f"mixture_score: {mixture_score}\n") + mixture_id = " ".join([learnware.id for learnware in mixture_learnware_list]) + print(f"mixture_learnware: {mixture_id}\n") + + rmtree(test_folder) # rm -r test_folder + + def test_learnware_reuse(self, learnware_num=5): + easy_market = self.test_upload_delete_learnware(learnware_num, delete=False) + print("Total Item:", len(easy_market)) + + 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) + user_info = BaseUserInfo(semantic_spec=user_semantic, stat_info={"RKMEStatSpecification": stat_spec}) + + _, _, _, mixture_learnware_list = easy_market.search_learnware(user_info) + + # Based on user information, the learnware market returns a list of learnwares (learnware_list) + # Use jobselector reuser to reuse the searched learnwares to make prediction + reuse_job_selector = JobSelectorReuser(learnware_list=mixture_learnware_list) + job_selector_predict_y = reuse_job_selector.predict(user_data=data_X) + + # Use averaging ensemble reuser to reuse the searched learnwares to make prediction + reuse_ensemble = AveragingReuser(learnware_list=mixture_learnware_list, mode="vote_by_prob") + ensemble_predict_y = reuse_ensemble.predict(user_data=data_X) + + # Use ensemble pruning reuser to reuse the searched learnwares to make prediction + reuse_ensemble = EnsemblePruningReuser(learnware_list=mixture_learnware_list, mode="classification") + reuse_ensemble.fit(train_X[-200:], train_y[-200:]) + ensemble_pruning_predict_y = reuse_ensemble.predict(user_data=data_X) + + print("Job Selector Acc:", np.sum(np.argmax(job_selector_predict_y, axis=1) == data_y) / len(data_y)) + print("Averaging Reuser Acc:", np.sum(np.argmax(ensemble_predict_y, axis=1) == data_y) / len(data_y)) + print("Ensemble Pruning Reuser Acc:", np.sum(ensemble_pruning_predict_y == data_y) / len(data_y)) + + +def suite(): + _suite = unittest.TestSuite() + _suite.addTest(TestAllWorkflow("test_prepare_learnware_randomly")) + _suite.addTest(TestAllWorkflow("test_upload_delete_learnware")) + _suite.addTest(TestAllWorkflow("test_search_semantics")) + _suite.addTest(TestAllWorkflow("test_stat_search")) + _suite.addTest(TestAllWorkflow("test_learnware_reuse")) + return _suite + + +if __name__ == "__main__": + runner = unittest.TextTestRunner() + runner.run(suite()) From d0dc4a4e8f6e99a3891540b2f3515120c1c47eed Mon Sep 17 00:00:00 2001 From: bxdd Date: Sat, 28 Oct 2023 14:43:35 +0800 Subject: [PATCH 17/35] [MNT] add new market test --- learnware/market/module.py | 3 ++- tests/test_market/test_easy.py | 4 ++-- 2 files changed, 4 insertions(+), 3 deletions(-) diff --git a/learnware/market/module.py b/learnware/market/module.py index 80da5d6..221650d 100644 --- a/learnware/market/module.py +++ b/learnware/market/module.py @@ -10,10 +10,11 @@ MARKET_CONFIG = { } -def instatiate_learnware_market(market_id, name="easy"): +def instatiate_learnware_market(market_id, name="easy", **kwargs): return LearnwareMarket( market_id=market_id, organizer=MARKET_CONFIG[name]["organizer"], checker=MARKET_CONFIG[name]["checker"], searcher=MARKET_CONFIG[name]["searcher"], + **kwargs ) diff --git a/tests/test_market/test_easy.py b/tests/test_market/test_easy.py index 38459ad..a0b4dbe 100644 --- a/tests/test_market/test_easy.py +++ b/tests/test_market/test_easy.py @@ -11,7 +11,7 @@ from sklearn.model_selection import train_test_split from shutil import copyfile, rmtree import learnware -from learnware.market import EasyMarket, BaseUserInfo +from learnware.market import instatiate_learnware_market, BaseUserInfo from learnware.learnware import JobSelectorReuser, AveragingReuser, EnsemblePruningReuser import learnware.specification as specification @@ -38,7 +38,7 @@ class TestAllWorkflow(unittest.TestCase): def _init_learnware_market(self): """initialize learnware market""" - easy_market = EasyMarket(market_id="sklearn_digits", rebuild=True) + easy_market = instatiate_learnware_market(market_id="sklearn_digits", name='easy', rebuild=True) return easy_market def test_prepare_learnware_randomly(self, learnware_num=5): From 88ef74b6fe864196564a755f62c8b0f1e6b94b5f Mon Sep 17 00:00:00 2001 From: bxdd Date: Sat, 28 Oct 2023 14:50:58 +0800 Subject: [PATCH 18/35] [FIX] fix bugs, ad tests --- learnware/market/base.py | 12 +++-- learnware/market/easy2/organizer.py | 6 ++- learnware/market/easy2/searcher.py | 2 +- tests/test_market/learnware_example/README.md | 10 ++++ .../learnware_example/environment.yaml | 27 +++++++++++ .../learnware_example/example.yaml | 8 ++++ .../learnware_example/example_init.py | 20 ++++++++ tests/test_market/test_easy.py | 48 ++++--------------- 8 files changed, 87 insertions(+), 46 deletions(-) create mode 100644 tests/test_market/learnware_example/README.md create mode 100644 tests/test_market/learnware_example/environment.yaml create mode 100644 tests/test_market/learnware_example/example.yaml create mode 100644 tests/test_market/learnware_example/example_init.py diff --git a/learnware/market/base.py b/learnware/market/base.py index 48dfd6b..f7673df 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -94,12 +94,15 @@ class LearnwareMarket: def get_learnware_by_ids(self, id: Union[str, List[str]], **kwargs) -> Union[Learnware, List[Learnware]]: return self.learnware_organizer.get_learnware_by_ids(id, **kwargs) + def __len__(self): + return len(self.learnware_organizer) + class LearnwareOrganizer: - def __init__(self, market_id=None, checker: 'LearnwareChecker' = None): + def __init__(self, market_id=None, checker: "LearnwareChecker" = None): self.reset(market_id=market_id, checker=checker) - - def reset(self, market_id=None, checker: 'LearnwareChecker'=None, **kwargs): + + def reset(self, market_id=None, checker: "LearnwareChecker" = None, **kwargs): self.market_id = market_id self.checker = checker @@ -242,6 +245,9 @@ class LearnwareOrganizer: """ raise NotImplementedError("get_learnwares is not implemented") + def __len__(self): + raise NotImplementedError("__len__ is not implemented") + class LearnwareSearcher: def __init__(self, organizer: LearnwareOrganizer = None): diff --git a/learnware/market/easy2/organizer.py b/learnware/market/easy2/organizer.py index 09a0f3d..c604006 100644 --- a/learnware/market/easy2/organizer.py +++ b/learnware/market/easy2/organizer.py @@ -30,7 +30,6 @@ logger = get_module_logger("easy_organizer") class EasyOrganizer(LearnwareOrganizer): - def reload_market(self, rebuild=False) -> bool: """Reload the learnware organizer when server restared. @@ -158,7 +157,7 @@ class EasyOrganizer(LearnwareOrganizer): if new_learnware is None: return None, EasyChecker.INVALID_LEARNWARE - learnwere_status = check_status if check_status is not None else self.checker.check_learnware(new_learnware) + learnwere_status = check_status if check_status is not None else self.checker(new_learnware) self.dbops.add_learnware( id=id, @@ -348,3 +347,6 @@ class EasyOrganizer(LearnwareOrganizer): return list(self.learnware_list.values()) else: return list(self.learnware_list.values())[:top] + + def __len__(self): + return len(self.learnware_list) diff --git a/learnware/market/easy2/searcher.py b/learnware/market/easy2/searcher.py index 2a64f9d..aa6388c 100644 --- a/learnware/market/easy2/searcher.py +++ b/learnware/market/easy2/searcher.py @@ -569,7 +569,7 @@ class EasySearcher(LearnwareSearcher): the third is the score of Learnware (mixture) the fourth is the list of Learnware (mixture), the size is search_num """ - learnware_list = [self.learnware_list[key] for key in self.learnware_list] + learnware_list = self.learnware_oganizer.get_learnwares() # learnware_list = self._search_by_semantic_spec_exact(learnware_list, user_info) # if len(learnware_list) == 0: learnware_list = self._search_by_semantic_spec_fuzz(learnware_list, user_info) diff --git a/tests/test_market/learnware_example/README.md b/tests/test_market/learnware_example/README.md new file mode 100644 index 0000000..51aac5a --- /dev/null +++ b/tests/test_market/learnware_example/README.md @@ -0,0 +1,10 @@ +## How to Generate Environment Yaml + +* create env config for conda: +```shell +conda env export | grep -v "^prefix: " > environment.yml +``` +* recover env from config +``` +conda env create -f environment.yml +``` \ No newline at end of file diff --git a/tests/test_market/learnware_example/environment.yaml b/tests/test_market/learnware_example/environment.yaml new file mode 100644 index 0000000..2923bdb --- /dev/null +++ b/tests/test_market/learnware_example/environment.yaml @@ -0,0 +1,27 @@ +name: learnware_example_env +channels: + - defaults +dependencies: + - _libgcc_mutex=0.1=main + - _openmp_mutex=5.1=1_gnu + - ca-certificates=2023.01.10=h06a4308_0 + - ld_impl_linux-64=2.38=h1181459_1 + - libffi=3.4.2=h6a678d5_6 + - libgcc-ng=11.2.0=h1234567_1 + - libgomp=11.2.0=h1234567_1 + - libstdcxx-ng=11.2.0=h1234567_1 + - ncurses=6.4=h6a678d5_0 + - openssl=1.1.1t=h7f8727e_0 + - pip=23.0.1=py38h06a4308_0 + - python=3.8.16=h7a1cb2a_3 + - readline=8.2=h5eee18b_0 + - setuptools=66.0.0=py38h06a4308_0 + - sqlite=3.41.2=h5eee18b_0 + - tk=8.6.12=h1ccaba5_0 + - wheel=0.38.4=py38h06a4308_0 + - xz=5.2.10=h5eee18b_1 + - zlib=1.2.13=h5eee18b_0 + - pip: + - joblib==1.2.0 + - learnware==0.0.1.99 + - numpy==1.19.5 diff --git a/tests/test_market/learnware_example/example.yaml b/tests/test_market/learnware_example/example.yaml new file mode 100644 index 0000000..254bca4 --- /dev/null +++ b/tests/test_market/learnware_example/example.yaml @@ -0,0 +1,8 @@ +model: + class_name: SVM + kwargs: {} +stat_specifications: + - module_path: learnware.specification + class_name: RKMEStatSpecification + file_name: svm.json + kwargs: {} \ No newline at end of file diff --git a/tests/test_market/learnware_example/example_init.py b/tests/test_market/learnware_example/example_init.py new file mode 100644 index 0000000..47d3708 --- /dev/null +++ b/tests/test_market/learnware_example/example_init.py @@ -0,0 +1,20 @@ +import os +import joblib +import numpy as np +from learnware.model import BaseModel + + +class SVM(BaseModel): + def __init__(self): + super(SVM, self).__init__(input_shape=(64,), output_shape=(10,)) + dir_path = os.path.dirname(os.path.abspath(__file__)) + self.model = joblib.load(os.path.join(dir_path, "svm.pkl")) + + def fit(self, X: np.ndarray, y: np.ndarray): + pass + + def predict(self, X: np.ndarray) -> np.ndarray: + return self.model.predict_proba(X) + + def finetune(self, X: np.ndarray, y: np.ndarray): + pass diff --git a/tests/test_market/test_easy.py b/tests/test_market/test_easy.py index a0b4dbe..2aa674c 100644 --- a/tests/test_market/test_easy.py +++ b/tests/test_market/test_easy.py @@ -12,7 +12,6 @@ from shutil import copyfile, rmtree import learnware from learnware.market import instatiate_learnware_market, BaseUserInfo -from learnware.learnware import JobSelectorReuser, AveragingReuser, EnsemblePruningReuser import learnware.specification as specification curr_root = os.path.dirname(os.path.abspath(__file__)) @@ -30,7 +29,7 @@ user_semantic = { } -class TestAllWorkflow(unittest.TestCase): +class TestMarket(unittest.TestCase): @classmethod def setUpClass(cls) -> None: np.random.seed(2023) @@ -38,7 +37,7 @@ class TestAllWorkflow(unittest.TestCase): def _init_learnware_market(self): """initialize learnware market""" - easy_market = instatiate_learnware_market(market_id="sklearn_digits", name='easy', rebuild=True) + easy_market = instatiate_learnware_market(market_id="sklearn_digits", name="easy", rebuild=True) return easy_market def test_prepare_learnware_randomly(self, learnware_num=5): @@ -100,13 +99,13 @@ class TestAllWorkflow(unittest.TestCase): easy_market.add_learnware(zip_path, semantic_spec) print("Total Item:", len(easy_market)) - curr_inds = easy_market._get_ids() + curr_inds = easy_market.get_learnware_ids() print("Available ids After Uploading Learnwares:", curr_inds) if delete: for learnware_id in curr_inds: easy_market.delete_learnware(learnware_id) - curr_inds = easy_market._get_ids() + curr_inds = easy_market.get_learnware_ids() print("Available ids After Deleting Learnwares:", curr_inds) return easy_market @@ -174,44 +173,13 @@ class TestAllWorkflow(unittest.TestCase): rmtree(test_folder) # rm -r test_folder - def test_learnware_reuse(self, learnware_num=5): - easy_market = self.test_upload_delete_learnware(learnware_num, delete=False) - print("Total Item:", len(easy_market)) - - 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) - user_info = BaseUserInfo(semantic_spec=user_semantic, stat_info={"RKMEStatSpecification": stat_spec}) - - _, _, _, mixture_learnware_list = easy_market.search_learnware(user_info) - - # Based on user information, the learnware market returns a list of learnwares (learnware_list) - # Use jobselector reuser to reuse the searched learnwares to make prediction - reuse_job_selector = JobSelectorReuser(learnware_list=mixture_learnware_list) - job_selector_predict_y = reuse_job_selector.predict(user_data=data_X) - - # Use averaging ensemble reuser to reuse the searched learnwares to make prediction - reuse_ensemble = AveragingReuser(learnware_list=mixture_learnware_list, mode="vote_by_prob") - ensemble_predict_y = reuse_ensemble.predict(user_data=data_X) - - # Use ensemble pruning reuser to reuse the searched learnwares to make prediction - reuse_ensemble = EnsemblePruningReuser(learnware_list=mixture_learnware_list, mode="classification") - reuse_ensemble.fit(train_X[-200:], train_y[-200:]) - ensemble_pruning_predict_y = reuse_ensemble.predict(user_data=data_X) - - print("Job Selector Acc:", np.sum(np.argmax(job_selector_predict_y, axis=1) == data_y) / len(data_y)) - print("Averaging Reuser Acc:", np.sum(np.argmax(ensemble_predict_y, axis=1) == data_y) / len(data_y)) - print("Ensemble Pruning Reuser Acc:", np.sum(ensemble_pruning_predict_y == data_y) / len(data_y)) - def suite(): _suite = unittest.TestSuite() - _suite.addTest(TestAllWorkflow("test_prepare_learnware_randomly")) - _suite.addTest(TestAllWorkflow("test_upload_delete_learnware")) - _suite.addTest(TestAllWorkflow("test_search_semantics")) - _suite.addTest(TestAllWorkflow("test_stat_search")) - _suite.addTest(TestAllWorkflow("test_learnware_reuse")) + _suite.addTest(TestMarket("test_prepare_learnware_randomly")) + _suite.addTest(TestMarket("test_upload_delete_learnware")) + _suite.addTest(TestMarket("test_search_semantics")) + _suite.addTest(TestMarket("test_stat_search")) return _suite From 9ed25fc6cd9a740bd1bbd5dac6820c6df5fc6adf Mon Sep 17 00:00:00 2001 From: Gene Date: Sat, 28 Oct 2023 16:47:37 +0800 Subject: [PATCH 19/35] [MNT] switch to BaseChecker, BaseOrganizer and BaseSearcher --- learnware/market/__init__.py | 2 +- learnware/market/base.py | 48 ++++++++++++++--------------- learnware/market/easy2/checker.py | 4 +-- learnware/market/easy2/organizer.py | 4 +-- learnware/market/easy2/searcher.py | 4 +-- 5 files changed, 31 insertions(+), 31 deletions(-) diff --git a/learnware/market/__init__.py b/learnware/market/__init__.py index 81a4184..bd939f5 100644 --- a/learnware/market/__init__.py +++ b/learnware/market/__init__.py @@ -1,5 +1,5 @@ from .anchor import AnchoredUserInfo, AnchoredOrganizer -from .base import BaseUserInfo, LearnwareMarket, LearnwareChecker, LearnwareOrganizer +from .base import BaseUserInfo, LearnwareMarket, BaseChecker, BaseOrganizer, BaseSearcher from .evolve_anchor import EvolveAnchoredOrganizer from .evolve import EvolvedOrganizer from .easy2 import EasyChecker, EasyOrganizer, EasySearcher diff --git a/learnware/market/base.py b/learnware/market/base.py index f7673df..dc91a9f 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -50,18 +50,18 @@ class LearnwareMarket: def __init__( self, market_id: str = None, - organizer: "LearnwareOrganizer" = None, - checker: "LearnwareChecker" = None, - searcher: "LearnwareSearcher" = None, + organizer: "BaseOrganizer" = None, + checker: "BaseChecker" = None, + searcher: "BaseSearcher" = None, rebuild=False, ): self.market_id = market_id - self.learnware_organizer = LearnwareOrganizer() if organizer is None else organizer - self.learnware_checker = LearnwareChecker() if checker is None else checker + self.learnware_organizer = BaseOrganizer() if organizer is None else organizer + self.learnware_checker = BaseChecker() if checker is None else checker self.learnware_checker.reset(organizer=self.learnware_organizer) self.learnware_organizer.reset(market_id=market_id, checker=self.learnware_checker) self.learnware_organizer.reload_market(rebuild=rebuild) - self.learnware_searcher = LearnwareSearcher() if searcher is None else searcher + self.learnware_searcher = BaseSearcher() if searcher is None else searcher self.learnware_searcher.reset(organizer=self.learnware_organizer) def reload_market(self, **kwargs) -> bool: @@ -98,11 +98,11 @@ class LearnwareMarket: return len(self.learnware_organizer) -class LearnwareOrganizer: - def __init__(self, market_id=None, checker: "LearnwareChecker" = None): +class BaseOrganizer: + def __init__(self, market_id=None, checker: BaseChecker = None): self.reset(market_id=market_id, checker=checker) - def reset(self, market_id=None, checker: "LearnwareChecker" = None, **kwargs): + def reset(self, market_id=None, checker: BaseChecker = None, **kwargs): self.market_id = market_id self.checker = checker @@ -115,7 +115,7 @@ class LearnwareOrganizer: A flag indicating whether the market is reload successfully. """ - raise NotImplementedError("reload market is Not Implemented") + raise NotImplementedError("reload market is Not Implemented in BaseOrganizer") def add_learnware(self, zip_path: str, semantic_spec: dict) -> Tuple[str, bool]: """Add a learnware into the market. @@ -145,7 +145,7 @@ class LearnwareOrganizer: file for model or statistical specification not found """ - raise NotImplementedError("add learnware is Not Implemented") + raise NotImplementedError("add learnware is Not Implemented in BaseOrganizer") def delete_learnware(self, id: str) -> bool: """Delete a learnware from market @@ -165,7 +165,7 @@ class LearnwareOrganizer: Exception Raise an excpetion when given id is NOT found in learnware list """ - raise NotImplementedError("delete learnware is Not Implemented") + raise NotImplementedError("delete learnware is Not Implemented in BaseOrganizer") def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, **kwargs) -> bool: """ @@ -176,7 +176,7 @@ class LearnwareOrganizer: id : str id of target learnware. """ - raise NotImplementedError("update learnware is Not Implemented") + raise NotImplementedError("update learnware is Not Implemented in BaseOrganizer") def get_learnware_by_ids(self, id: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """ @@ -195,7 +195,7 @@ class LearnwareOrganizer: - The returned items are search results. - 'None' indicating the target id not found. """ - raise NotImplementedError("get_learnware_by_ids is not implemented") + raise NotImplementedError("get_learnware_by_ids is not implemented in BaseOrganizer") def get_learnware_path_by_ids(self, ids: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """Get Zipped Learnware file by id @@ -213,7 +213,7 @@ class LearnwareOrganizer: Return the path for target learnware or list of path. None for Learnware NOT Found. """ - raise NotImplementedError("get_learnware_path_by_ids is not implemented") + raise NotImplementedError("get_learnware_path_by_ids is not implemented in BaseOrganizer") def get_learnware_ids(self, top: int = None) -> List[str]: """get the list of learnware ids @@ -228,7 +228,7 @@ class LearnwareOrganizer: List[str] the first top ids """ - raise NotImplementedError("get_learnware_ids is not implemented") + raise NotImplementedError("get_learnware_ids is not implemented in BaseOrganizer") def get_learnwares(self, top: int = None) -> List[Learnware]: """get the list of learnwares @@ -243,14 +243,14 @@ class LearnwareOrganizer: List[Learnware] the first top learnwares """ - raise NotImplementedError("get_learnwares is not implemented") + raise NotImplementedError("get_learnwares is not implemented in BaseOrganizer") def __len__(self): - raise NotImplementedError("__len__ is not implemented") + raise NotImplementedError("__len__ is not implemented in BaseOrganizer") -class LearnwareSearcher: - def __init__(self, organizer: LearnwareOrganizer = None): +class BaseSearcher: + def __init__(self, organizer: BaseOrganizer = None): self.learnware_oganizer = organizer def reset(self, organizer): @@ -264,15 +264,15 @@ class LearnwareSearcher: user_info : BaseUserInfo user_info contains semantic_spec and stat_info """ - raise NotImplementedError("'__call__' method is not implemented in LearnwareSearcher") + raise NotImplementedError("'__call__' method is not implemented in BaseSearcher") -class LearnwareChecker: +class BaseChecker: INVALID_LEARNWARE = -1 NONUSABLE_LEARNWARE = 0 USABLE_LEARWARE = 1 - def __init__(self, organizer: LearnwareOrganizer = None): + def __init__(self, organizer: BaseOrganizer = None): self.learnware_oganizer = organizer def reset(self, organizer): @@ -294,4 +294,4 @@ class LearnwareChecker: - The NOPREDICTION_LEARNWARE denotes the leanrware pass the check and can make prediction """ - raise NotImplementedError("'__call__' method is not implemented in LearnwareChecker") + raise NotImplementedError("'__call__' method is not implemented in BaseChecker") diff --git a/learnware/market/easy2/checker.py b/learnware/market/easy2/checker.py index 8a4a250..062f9ad 100644 --- a/learnware/market/easy2/checker.py +++ b/learnware/market/easy2/checker.py @@ -2,13 +2,13 @@ import traceback import numpy as np import torch -from ..base import LearnwareChecker +from ..base import BaseChecker from ...logger import get_module_logger logger = get_module_logger("easy_checker", "INFO") -class EasyChecker(LearnwareChecker): +class EasyChecker(BaseChecker): def __call__(self, learnware): semantic_spec = learnware.get_specification().get_semantic_spec() diff --git a/learnware/market/easy2/organizer.py b/learnware/market/easy2/organizer.py index c604006..2f2c062 100644 --- a/learnware/market/easy2/organizer.py +++ b/learnware/market/easy2/organizer.py @@ -23,13 +23,13 @@ from ...logger import get_module_logger from ...learnware import Learnware, get_learnware_from_dirpath from ...specification import RKMEStatSpecification, Specification -from ..base import LearnwareOrganizer, LearnwareChecker +from ..base import BaseOrganizer, BaseChecker from ...logger import get_module_logger logger = get_module_logger("easy_organizer") -class EasyOrganizer(LearnwareOrganizer): +class EasyOrganizer(BaseOrganizer): def reload_market(self, rebuild=False) -> bool: """Reload the learnware organizer when server restared. diff --git a/learnware/market/easy2/searcher.py b/learnware/market/easy2/searcher.py index aa6388c..cd2759f 100644 --- a/learnware/market/easy2/searcher.py +++ b/learnware/market/easy2/searcher.py @@ -4,7 +4,7 @@ from rapidfuzz import fuzz from cvxopt import solvers, matrix from typing import Tuple, List -from ..base import BaseUserInfo, LearnwareSearcher +from ..base import BaseUserInfo, BaseSearcher from ...learnware import Learnware from ...specification import RKMEStatSpecification from ...logger import get_module_logger @@ -12,7 +12,7 @@ from ...logger import get_module_logger logger = get_module_logger("easy_seacher") -class EasySearcher(LearnwareSearcher): +class EasySearcher(BaseSearcher): def _convert_dist_to_score( self, dist_list: List[float], dist_epsilon: float = 0.01, min_score: float = 0.92 ) -> List[float]: From 087659748270fa5eee34cd7c6097c91ac7a34d82 Mon Sep 17 00:00:00 2001 From: Gene Date: Sat, 28 Oct 2023 21:44:48 +0800 Subject: [PATCH 20/35] [MNT] change single checker to multiple checker --- learnware/market/__init__.py | 2 +- learnware/market/base.py | 67 +++++++++++---- learnware/market/easy2/__init__.py | 2 +- learnware/market/easy2/checker.py | 94 +++++++++++++++------ learnware/market/easy2/organizer.py | 51 ++--------- learnware/market/evolve_anchor/organizer.py | 4 +- learnware/market/module.py | 6 +- 7 files changed, 131 insertions(+), 95 deletions(-) diff --git a/learnware/market/__init__.py b/learnware/market/__init__.py index bd939f5..bc6137e 100644 --- a/learnware/market/__init__.py +++ b/learnware/market/__init__.py @@ -2,7 +2,7 @@ from .anchor import AnchoredUserInfo, AnchoredOrganizer from .base import BaseUserInfo, LearnwareMarket, BaseChecker, BaseOrganizer, BaseSearcher from .evolve_anchor import EvolveAnchoredOrganizer from .evolve import EvolvedOrganizer -from .easy2 import EasyChecker, EasyOrganizer, EasySearcher +from .easy2 import EasyOrganizer, EasySearcher, EasySemanticChecker, EasyStatisticalChecker from .hetergeneous import HeterogeneousOrganizer, MappingFunction from .easy import EasyMarket diff --git a/learnware/market/base.py b/learnware/market/base.py index dc91a9f..1bc2d2d 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -1,11 +1,12 @@ import os import torch +import tempfile import traceback import numpy as np from typing import Tuple, Any, List, Union -from ..learnware import Learnware +from ..learnware import Learnware, get_learnware_from_dirpath from ..logger import get_module_logger logger = get_module_logger("market_base", "INFO") @@ -51,27 +52,57 @@ class LearnwareMarket: self, market_id: str = None, organizer: "BaseOrganizer" = None, - checker: "BaseChecker" = None, searcher: "BaseSearcher" = None, + checker_list: List["BaseChecker"] = None, rebuild=False, ): self.market_id = market_id self.learnware_organizer = BaseOrganizer() if organizer is None else organizer - self.learnware_checker = BaseChecker() if checker is None else checker - self.learnware_checker.reset(organizer=self.learnware_organizer) - self.learnware_organizer.reset(market_id=market_id, checker=self.learnware_checker) + self.learnware_organizer.reset(market_id=market_id) self.learnware_organizer.reload_market(rebuild=rebuild) self.learnware_searcher = BaseSearcher() if searcher is None else searcher self.learnware_searcher.reset(organizer=self.learnware_organizer) + + if checker_list is None: + self.learnware_checker = {"BaseChecker": BaseChecker()} + else: + self.learnware_checker = {checker.__class__.__name__: checker for checker in checker_list} + for name, checker in self.learnware_checker.items(): + checker.reset(organizer=self.learnware_organizer) def reload_market(self, **kwargs) -> bool: self.learnware_organizer.reload_market(**kwargs) - def check_learnware(self, learnware: Learnware, **kwargs) -> bool: - return self.learnware_checker(learnware, **kwargs) - - def add_learnware(self, zip_path: str, semantic_spec: dict, **kwargs) -> Tuple[str, bool]: - return self.learnware_organizer.add_learnware(zip_path, semantic_spec, **kwargs) + def check_learnware(self, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs) -> bool: + try: + with tempfile.TemporaryDirectory(prefix="pending_learnware_") as tempdir: + with zipfile.ZipFile(zip_path, mode="r") as z_file: + z_file.extractall(tempdir) + + pending_learnware = get_learnware_from_dirpath( + id="pending", semantic_spec=semantic_specification, learnware_dirpath=tempdir + ) + + final_status = BaseChecker.INVALID_LEARNWARE + checker_names = list(self.learnware_checker.keys()) if checker_names is None else checker_names + + for name in checker_names: + checker = self.learnware_checker[name] + check_status = checker(pending_learnware) + final_status = max(final_status, check_status) + + if check_status == BaseChecker.INVALID_LEARNWARE: + return BaseChecker.INVALID_LEARNWARE + + return final_status + + except Exception as err: + logger.warning(f"Check learnware failed! Due to {err}.") + return BaseChecker.INVALID_LEARNWARE + + def add_learnware(self, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs) -> Tuple[str, bool]: + check_status = self.check_learnware(zip_path, semantic_spec, checker_names) + return self.learnware_organizer.add_learnware(zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs) def search_learnware(self, user_info: BaseUserInfo, **kwargs) -> Tuple[Any, List[Learnware]]: return self.learnware_searcher(user_info, **kwargs) @@ -79,8 +110,9 @@ class LearnwareMarket: def delete_learnware(self, id: str, **kwargs) -> bool: return self.learnware_organizer.delete_learnware(id, **kwargs) - def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, **kwargs) -> bool: - return self.learnware_organizer.update_learnware(id, zip_path=zip_path, semantic_spec=semantic_spec, **kwargs) + def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs) -> bool: + check_status = self.check_learnware(zip_path, semantic_spec, checker_names) + return self.learnware_organizer.update_learnware(id, zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs) def get_learnware_ids(self, top: int = None, **kwargs): return self.learnware_organizer.get_learnware_ids(top, **kwargs) @@ -99,12 +131,11 @@ class LearnwareMarket: class BaseOrganizer: - def __init__(self, market_id=None, checker: BaseChecker = None): - self.reset(market_id=market_id, checker=checker) + def __init__(self, market_id=None): + self.reset(market_id=market_id) - def reset(self, market_id=None, checker: BaseChecker = None, **kwargs): + def reset(self, market_id=None, **kwargs): self.market_id = market_id - self.checker = checker def reload_market(self, rebuild=False, **kwargs) -> bool: """Reload the learnware organizer when server restared. @@ -117,7 +148,7 @@ class BaseOrganizer: raise NotImplementedError("reload market is Not Implemented in BaseOrganizer") - def add_learnware(self, zip_path: str, semantic_spec: dict) -> Tuple[str, bool]: + def add_learnware(self, zip_path: str, semantic_spec: dict, check_status: int) -> Tuple[str, bool]: """Add a learnware into the market. .. note:: @@ -167,7 +198,7 @@ class BaseOrganizer: """ raise NotImplementedError("delete learnware is Not Implemented in BaseOrganizer") - def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, **kwargs) -> bool: + def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, check_status: int) -> bool: """ Update Learnware with id and content to be updated. diff --git a/learnware/market/easy2/__init__.py b/learnware/market/easy2/__init__.py index 2ab8c48..2178119 100644 --- a/learnware/market/easy2/__init__.py +++ b/learnware/market/easy2/__init__.py @@ -1,3 +1,3 @@ from .organizer import EasyOrganizer -from .checker import EasyChecker from .searcher import EasySearcher +from .checker import EasySemanticChecker, EasyStatisticalChecker diff --git a/learnware/market/easy2/checker.py b/learnware/market/easy2/checker.py index 062f9ad..25ee452 100644 --- a/learnware/market/easy2/checker.py +++ b/learnware/market/easy2/checker.py @@ -3,71 +3,113 @@ import numpy as np import torch from ..base import BaseChecker +from ...config import C from ...logger import get_module_logger logger = get_module_logger("easy_checker", "INFO") -class EasyChecker(BaseChecker): +class EasySemanticChecker(BaseChecker): + def __call__(self, learnware): + semantic_spec = learnware.get_specification().get_semantic_spec() + try: + for key in C["semantic_specs"]: + value = semantic_spec[key]["Values"] + valid_type = C["semantic_specs"][key]["Type"] + assert semantic_spec[key]["Type"] == valid_type, f"{key} type mismatch" + + if valid_type == "Class": + valid_list = C["semantic_specs"][key]["Values"] + assert len(value) == 1, f"{key} must be unique" + assert value[0] in valid_list, f"{key} must be in {valid_list}" + + elif valid_type == "Tag": + valid_list = C["semantic_specs"][key]["Values"] + assert len(value) >= 1, f"{key} cannot be empty" + for v in value: + assert v in valid_list, f"{key} must be in {valid_list}" + + elif valid_type == "String": + assert isinstance(value, str), f"{key} must be string" + assert len(value) >= 1, f"{key} cannot be empty" + + if semantic_spec["Data"]["Values"][0] == "Table": + assert semantic_spec["Input"] is not None, "Lack of input semantics" + dim = semantic_spec["Input"]["Dimension"] + for k, v in semantic_spec["Input"]["Description"].items(): + assert int(k) >= 0 and int(k) < dim, f"Dimension number in [0, {dim})" + assert isinstance(v, str), "Description must be string" + + if semantic_spec["Task"]["Values"][0] in ["Classification", "Regression", "Feature Extraction"]: + assert semantic_spec["Output"] is not None, "Lack of output semantics" + dim = semantic_spec["Output"]["Dimension"] + for k, v in semantic_spec["Output"]["Description"].items(): + assert int(k) >= 0 and int(k) < dim, f"Dimension number in [0, {dim})" + assert isinstance(v, str), "Description must be string" + + return self.NONUSABLE_LEARNWARE + + except Exception as err: + logger.warning(f"semantic_specification is not valid due to {err}!") + return self.INVALID_LEARNWARE + + +class EasyStatisticalChecker(BaseChecker): def __call__(self, learnware): semantic_spec = learnware.get_specification().get_semantic_spec() try: - # check model instantiation + # Check model instantiation learnware.instantiate_model() except Exception as e: traceback.print_exc() - logger.warning(f"The learnware [{learnware.id}] is instantiated failed! Due to {e}") - return self.NONUSABLE_LEARNWARE + logger.warning(f"The learnware [{learnware.id}] is instantiated failed! Due to {e}.") + return self.INVALID_LEARNWARE try: learnware_model = learnware.get_model() - # check input shape + # Check input shape if semantic_spec["Data"]["Values"][0] == "Table": input_shape = (semantic_spec["Input"]["Dimension"],) else: input_shape = learnware_model.input_shape - pass - # check rkme dimension + # Check rkme dimension stat_spec = learnware.get_specification().get_stat_spec_by_name("RKMEStatSpecification") if stat_spec is not None: if stat_spec.get_z().shape[1:] != input_shape: - logger.warning(f"The learnware [{learnware.id}] input dimension mismatch with stat specification") - return self.NONUSABLE_LEARNWARE - pass + logger.warning(f"The learnware [{learnware.id}] input dimension mismatch with stat specification.") + return self.INVALID_LEARNWARE inputs = np.random.randn(10, *input_shape) outputs = learnware.predict(inputs) - # check output + # Check output if outputs.ndim == 1: outputs = outputs.reshape(-1, 1) - pass + + if outputs.shape[1:] != learnware_model.output_shape: + logger.warning(f"The learnware [{learnware.id}] output dimention mismatch!") + return self.INVALID_LEARNWARE if semantic_spec["Task"]["Values"][0] in ("Classification", "Regression", "Feature Extraction"): - # check output type + # Check output type if isinstance(outputs, torch.Tensor): outputs = outputs.detach().cpu().numpy() if not isinstance(outputs, np.ndarray): - logger.warning(f"The learnware [{learnware.id}] output must be np.ndarray or torch.Tensor") - return self.NONUSABLE_LEARNWARE + logger.warning(f"The learnware [{learnware.id}] output must be np.ndarray or torch.Tensor!") + return self.INVALID_LEARNWARE - # check output shape + # Check output shape output_dim = int(semantic_spec["Output"]["Dimension"]) if outputs[0].shape[0] != output_dim: - logger.warning(f"The learnware [{learnware.id}] input and output dimention is error") - return self.NONUSABLE_LEARNWARE - pass - else: - if outputs.shape[1:] != learnware_model.output_shape: - logger.warning(f"The learnware [{learnware.id}] input and output dimention is error") - return self.NONUSABLE_LEARNWARE + logger.warning(f"The learnware [{learnware.id}] output dimention mismatch!") + return self.INVALID_LEARNWARE except Exception as e: - logger.warning(f"The learnware [{learnware.id}] prediction is not avaliable! Due to {repr(e)}") - return self.NONUSABLE_LEARNWARE + logger.warning(f"The learnware [{learnware.id}] prediction is not avaliable! Due to {repr(e)}.") + return self.INVALID_LEARNWARE - return self.USABLE_LEARWARE + return self.USABLE_LEARWARE \ No newline at end of file diff --git a/learnware/market/easy2/organizer.py b/learnware/market/easy2/organizer.py index 2f2c062..55780e3 100644 --- a/learnware/market/easy2/organizer.py +++ b/learnware/market/easy2/organizer.py @@ -13,7 +13,6 @@ from shutil import copyfile, rmtree from typing import Tuple, Any, List, Union, Dict from .database_ops import DatabaseOperations -from .checker import EasyChecker from ..base import LearnwareMarket, BaseUserInfo @@ -95,42 +94,6 @@ class EasyOrganizer(BaseOrganizer): """ semantic_spec = copy.deepcopy(semantic_spec) - - if not os.path.exists(zip_path): - logger.warning("Zip Path NOT Found! Fail to add learnware.") - return None, EasyChecker.INVALID_LEARNWARE - - try: - if len(semantic_spec["Data"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please choose Data.") - return None, EasyChecker.INVALID_LEARNWARE - if len(semantic_spec["Task"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please choose Task.") - return None, EasyChecker.INVALID_LEARNWARE - if len(semantic_spec["Library"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please choose Device.") - return None, EasyChecker.INVALID_LEARNWARE - if len(semantic_spec["Name"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please provide Name.") - return None, EasyChecker.INVALID_LEARNWARE - if len(semantic_spec["Description"]["Values"]) == 0 and len(semantic_spec["Scenario"]["Values"]) == 0: - logger.warning("Illegal semantic specification, please provide Scenario or Description.") - return None, EasyChecker.INVALID_LEARNWARE - if ( - semantic_spec["Data"]["Type"] != "Class" - or semantic_spec["Task"]["Type"] != "Class" - or semantic_spec["Library"]["Type"] != "Class" - or semantic_spec["Scenario"]["Type"] != "Tag" - or semantic_spec["Name"]["Type"] != "String" - or semantic_spec["Description"]["Type"] != "String" - ): - logger.warning("Illegal semantic specification, please provide the right type.") - return None, EasyChecker.INVALID_LEARNWARE - except: - print(semantic_spec) - logger.warning("Illegal semantic specification, some keys are missing.") - return None, EasyChecker.INVALID_LEARNWARE - logger.info("Get new learnware from %s" % (zip_path)) id = id if id is not None else "%08d" % (self.count) @@ -152,12 +115,12 @@ class EasyOrganizer(BaseOrganizer): rmtree(target_folder_dir) except: pass - return None, EasyChecker.INVALID_LEARNWARE + return None, BaseChecker.INVALID_LEARNWARE if new_learnware is None: - return None, EasyChecker.INVALID_LEARNWARE + return None, BaseChecker.INVALID_LEARNWARE - learnwere_status = check_status if check_status is not None else self.checker(new_learnware) + learnwere_status = check_status if check_status is not None else BaseChecker.NONUSABLE_LEARNWARE self.dbops.add_learnware( id=id, @@ -227,7 +190,7 @@ class EasyOrganizer(BaseOrganizer): assert ( zip_path is None and semantic_spec is None ), f"at least one of 'zip_path' and 'semantic_spec' should not be None when update learnware" - assert check_status != EasyChecker.INVALID_LEARNWARE, f"'check_status' can not be INVALID_LEARNWARE" + assert check_status != BaseChecker.INVALID_LEARNWARE, f"'check_status' can not be INVALID_LEARNWARE" if zip_path is None and check_status is not None: logger.warning("check_status will be ignored when zip_path is None for learnware update") @@ -252,12 +215,12 @@ class EasyOrganizer(BaseOrganizer): id=id, semantic_spec=semantic_spec, learnware_dirpath=tempdir ) except Exception: - return EasyChecker.INVALID_LEARNWARE + return BaseChecker.INVALID_LEARNWARE if new_learnware is None: - return EasyChecker.INVALID_LEARNWARE + return BaseChecker.INVALID_LEARNWARE - learnwere_status = self.checker.check_learnware(new_learnware) + learnwere_status = BaseChecker.NONUSABLE_LEARNWARE else: learnwere_status = self.use_flags[id] if zip_path is None else check_status diff --git a/learnware/market/evolve_anchor/organizer.py b/learnware/market/evolve_anchor/organizer.py index 1e8173e..04e9779 100644 --- a/learnware/market/evolve_anchor/organizer.py +++ b/learnware/market/evolve_anchor/organizer.py @@ -1,7 +1,7 @@ from typing import List -from ..evolve.organizer import EvolvedOrganizer -from ..anchor.organizer import AnchoredOrganizer, AnchoredUserInfo +from ..evolve import EvolvedOrganizer +from ..anchor import AnchoredOrganizer, AnchoredUserInfo from ...logger import get_module_logger logger = get_module_logger("evolve_anchor_organizer") diff --git a/learnware/market/module.py b/learnware/market/module.py index 221650d..57821cd 100644 --- a/learnware/market/module.py +++ b/learnware/market/module.py @@ -1,11 +1,11 @@ from .base import LearnwareMarket -from .easy2 import EasyChecker, EasyOrganizer, EasySearcher +from .easy2 import EasyOrganizer, EasySearcher, EasySemanticChecker, EasyStatisticalChecker MARKET_CONFIG = { "easy": { "organizer": EasyOrganizer(), - "checker": EasyChecker(), "searcher": EasySearcher(), + "checker_list": [EasySemanticChecker(), EasyStatisticalChecker()], } } @@ -14,7 +14,7 @@ def instatiate_learnware_market(market_id, name="easy", **kwargs): return LearnwareMarket( market_id=market_id, organizer=MARKET_CONFIG[name]["organizer"], - checker=MARKET_CONFIG[name]["checker"], searcher=MARKET_CONFIG[name]["searcher"], + checker_list=MARKET_CONFIG[name]["checker_list"], **kwargs ) From 05005d528a88982b4bc620e1015ebbaece2ed246 Mon Sep 17 00:00:00 2001 From: Gene Date: Sat, 28 Oct 2023 21:45:51 +0800 Subject: [PATCH 21/35] [MNT] format code by black --- .../pfs/pfs_cross_transfer.py | 4 ++- learnware/market/base.py | 26 ++++++++++++------- learnware/market/easy2/checker.py | 10 +++---- 3 files changed, 25 insertions(+), 15 deletions(-) diff --git a/examples/dataset_pfs_workflow/pfs/pfs_cross_transfer.py b/examples/dataset_pfs_workflow/pfs/pfs_cross_transfer.py index 93a3fa3..5f69127 100644 --- a/examples/dataset_pfs_workflow/pfs/pfs_cross_transfer.py +++ b/examples/dataset_pfs_workflow/pfs/pfs_cross_transfer.py @@ -85,7 +85,9 @@ def get_split_errs(algo): split = train_xs.shape[0] - proportion_list[tmp] model.fit( - train_xs[split:,], + train_xs[ + split:, + ], train_ys[split:], eval_set=[(val_xs, val_ys)], early_stopping_rounds=50, diff --git a/learnware/market/base.py b/learnware/market/base.py index 1bc2d2d..1c26cea 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -62,7 +62,7 @@ class LearnwareMarket: self.learnware_organizer.reload_market(rebuild=rebuild) self.learnware_searcher = BaseSearcher() if searcher is None else searcher self.learnware_searcher.reset(organizer=self.learnware_organizer) - + if checker_list is None: self.learnware_checker = {"BaseChecker": BaseChecker()} else: @@ -78,11 +78,11 @@ class LearnwareMarket: with tempfile.TemporaryDirectory(prefix="pending_learnware_") as tempdir: with zipfile.ZipFile(zip_path, mode="r") as z_file: z_file.extractall(tempdir) - + pending_learnware = get_learnware_from_dirpath( id="pending", semantic_spec=semantic_specification, learnware_dirpath=tempdir ) - + final_status = BaseChecker.INVALID_LEARNWARE checker_names = list(self.learnware_checker.keys()) if checker_names is None else checker_names @@ -93,16 +93,20 @@ class LearnwareMarket: if check_status == BaseChecker.INVALID_LEARNWARE: return BaseChecker.INVALID_LEARNWARE - + return final_status - + except Exception as err: logger.warning(f"Check learnware failed! Due to {err}.") return BaseChecker.INVALID_LEARNWARE - def add_learnware(self, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs) -> Tuple[str, bool]: + def add_learnware( + self, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs + ) -> Tuple[str, bool]: check_status = self.check_learnware(zip_path, semantic_spec, checker_names) - return self.learnware_organizer.add_learnware(zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs) + return self.learnware_organizer.add_learnware( + zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs + ) def search_learnware(self, user_info: BaseUserInfo, **kwargs) -> Tuple[Any, List[Learnware]]: return self.learnware_searcher(user_info, **kwargs) @@ -110,9 +114,13 @@ class LearnwareMarket: def delete_learnware(self, id: str, **kwargs) -> bool: return self.learnware_organizer.delete_learnware(id, **kwargs) - def update_learnware(self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs) -> bool: + def update_learnware( + self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs + ) -> bool: check_status = self.check_learnware(zip_path, semantic_spec, checker_names) - return self.learnware_organizer.update_learnware(id, zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs) + return self.learnware_organizer.update_learnware( + id, zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs + ) def get_learnware_ids(self, top: int = None, **kwargs): return self.learnware_organizer.get_learnware_ids(top, **kwargs) diff --git a/learnware/market/easy2/checker.py b/learnware/market/easy2/checker.py index 25ee452..7f26b91 100644 --- a/learnware/market/easy2/checker.py +++ b/learnware/market/easy2/checker.py @@ -17,18 +17,18 @@ class EasySemanticChecker(BaseChecker): value = semantic_spec[key]["Values"] valid_type = C["semantic_specs"][key]["Type"] assert semantic_spec[key]["Type"] == valid_type, f"{key} type mismatch" - + if valid_type == "Class": valid_list = C["semantic_specs"][key]["Values"] assert len(value) == 1, f"{key} must be unique" assert value[0] in valid_list, f"{key} must be in {valid_list}" - + elif valid_type == "Tag": valid_list = C["semantic_specs"][key]["Values"] assert len(value) >= 1, f"{key} cannot be empty" for v in value: assert v in valid_list, f"{key} must be in {valid_list}" - + elif valid_type == "String": assert isinstance(value, str), f"{key} must be string" assert len(value) >= 1, f"{key} cannot be empty" @@ -89,7 +89,7 @@ class EasyStatisticalChecker(BaseChecker): # Check output if outputs.ndim == 1: outputs = outputs.reshape(-1, 1) - + if outputs.shape[1:] != learnware_model.output_shape: logger.warning(f"The learnware [{learnware.id}] output dimention mismatch!") return self.INVALID_LEARNWARE @@ -112,4 +112,4 @@ class EasyStatisticalChecker(BaseChecker): logger.warning(f"The learnware [{learnware.id}] prediction is not avaliable! Due to {repr(e)}.") return self.INVALID_LEARNWARE - return self.USABLE_LEARWARE \ No newline at end of file + return self.USABLE_LEARWARE From ee9655e8337f6b8d5f3e16d326eb995695db4de0 Mon Sep 17 00:00:00 2001 From: Gene Date: Sat, 28 Oct 2023 22:24:17 +0800 Subject: [PATCH 22/35] [MNT] modify AnchoredOrganizer --- learnware/market/anchor/__init__.py | 3 +- learnware/market/anchor/organizer.py | 49 +--------------------------- 2 files changed, 3 insertions(+), 49 deletions(-) diff --git a/learnware/market/anchor/__init__.py b/learnware/market/anchor/__init__.py index 64e0622..f2b453c 100644 --- a/learnware/market/anchor/__init__.py +++ b/learnware/market/anchor/__init__.py @@ -1 +1,2 @@ -from .organizer import AnchoredOrganizer, AnchoredUserInfo +from .organizer import AnchoredOrganizer +from .searcher import AnchoredUserInfo diff --git a/learnware/market/anchor/organizer.py b/learnware/market/anchor/organizer.py index f903f1f..e91f597 100644 --- a/learnware/market/anchor/organizer.py +++ b/learnware/market/anchor/organizer.py @@ -1,37 +1,11 @@ from typing import List, Dict, Tuple, Any -from ..base import BaseUserInfo from ..easy2.organizer import EasyOrganizer from ...logger import get_module_logger from ...learnware import Learnware from ...specification import BaseStatSpecification -logger = get_module_logger("evolve_organizer") - - -class AnchoredUserInfo(BaseUserInfo): - """ - User Information for searching learnware (add the anchor design) - - - UserInfo contains the anchor list acquired from the market - - UserInfo can update stat_info based on anchors - """ - - def __init__(self, id: str, semantic_spec: dict = None, stat_info: dict = None, anchor_scores: dict = None): - super(AnchoredUserInfo, self).__init__(id, semantic_spec, stat_info) - self.anchor_scores = {} if anchor_scores is None else anchor_scores - - def update_anchor_score(self, id: str, score): - """Update score of anchor learnwares - - Parameters - ---------- - id : str - id of anchor learnwares - score : Any - score of anchor learnwares - """ - self.anchor_scores[id] = score +logger = get_module_logger("anchor_organizer") class AnchoredOrganizer(EasyOrganizer): @@ -86,24 +60,3 @@ class AnchoredOrganizer(EasyOrganizer): Learnwares for updating anchor_learnware_list """ pass - - def search_learnware(self, user_info: AnchoredUserInfo, anchored: bool = False) -> Tuple[Any, List[Learnware]]: - """Search learnwares with anchor marget - - if 'anchor' == True, search anchor Learnwares from anchor_learnware_list based on user_info - - if 'anchor' == False, find helpful learnwares from learnware_list based on user_info - - Parameters - ---------- - user_info : AnchoredUserInfo - - user_info with semantic specifications and statistical information - - some statistical information calculated on anchor learnwares - - Returns - ------- - Tuple[Any, List[Any]] - return two items: - - - first is recommended combination, None when no recommended combination is calculated or statistical specification is not provided. - - second is a list of matched learnwares - """ - pass From 5f4a63aa2951216e72f27340fd94b31ec817a47d Mon Sep 17 00:00:00 2001 From: Gene Date: Sat, 28 Oct 2023 22:24:44 +0800 Subject: [PATCH 23/35] [ENH] add AnchoredSearcher --- learnware/market/anchor/searcher.py | 108 ++++++++++++++++++++++++++++ 1 file changed, 108 insertions(+) create mode 100644 learnware/market/anchor/searcher.py diff --git a/learnware/market/anchor/searcher.py b/learnware/market/anchor/searcher.py new file mode 100644 index 0000000..0ccdd34 --- /dev/null +++ b/learnware/market/anchor/searcher.py @@ -0,0 +1,108 @@ +from typing import List, Dict, Tuple, Any, Union + +from ..base import BaseUserInfo +from ..easy2.searcher import EasySearcher +from ...learnware import Learnware + + +class AnchoredUserInfo(BaseUserInfo): + """ + User Information for searching learnware (add the anchor design) + + - UserInfo contains the anchor id list acquired from the market + - UserInfo can update stat_info based on anchors + """ + + def __init__( + self, id: str, semantic_spec: dict = None, stat_info: dict = None, anchor_learnware_ids: List[str] = None + ): + super(AnchoredUserInfo, self).__init__(id, semantic_spec, stat_info) + self.anchor_learnware_ids = [] if anchor_learnware_ids is None else anchor_learnware_ids + + def add_anchor_learnware_ids(self, learnware_ids: Union[str, List[str]]): + """Add the anchor learnware ids acquired from the market + + Parameters + ---------- + learnware_ids : Union[str, List[str]] + Anchor learnware ids + """ + if isinstance(learnware_ids, str): + learnware_ids = [learnware_ids] + self.anchor_learnware_ids += learnware_ids + + def update_stat_info(self, name: str, item: Any): + """Update stat_info based on anchor learnwares + + Parameters + ---------- + name : str + Name of stat_info + item : Any + Statistical information calculated on anchor learnwares + """ + self.stat_info[name] = item + + +class AnchoredSearcher(EasySearcher): + def search_anchor_learnware(self, user_info: AnchoredUserInfo) -> Tuple[Any, List[Learnware]]: + """Search anchor Learnwares from anchor_learnware_list based on user_info + + Parameters + ---------- + user_info : AnchoredUserInfo + - user_info with semantic specifications and statistical information + - some statistical information calculated on previous anchor learnwares + + Returns + ------- + Tuple[Any, List[Learnware]]: + return two items: + + - first is the usage of anchor learnwares, e.g., how to use anchors to calculate some statistical information + - second is a list of anchor learnwares + """ + pass + + def search_learnware(self, user_info: AnchoredUserInfo) -> Tuple[Any, List[Learnware]]: + """Find helpful learnwares from learnware_list based on user_info + + Parameters + ---------- + user_info : AnchoredUserInfo + - user_info with semantic specifications and statistical information + - some statistical information calculated on anchor learnwares + + Returns + ------- + Tuple[Any, List[Any]] + return two items: + + - first is recommended combination, None when no recommended combination is calculated or statistical specification is not provided. + - second is a list of matched learnwares + """ + pass + + def __call__(self, user_info: AnchoredUserInfo, anchor_flag: bool = False) -> Tuple[Any, List[Learnware]]: + """Search learnwares with anchor marget + - if 'anchor_flag' == True, search anchor Learnwares from anchor_learnware_list based on user_info + - if 'anchor_flag' == False, find helpful learnwares from learnware_list based on user_info + + Parameters + ---------- + user_info : AnchoredUserInfo + - user_info with semantic specifications and statistical information + - some statistical information calculated on anchor learnwares + + Returns + ------- + Tuple[Any, List[Any]] + return two items: + + - first is recommended combination, None when no recommended combination is calculated or statistical specification is not provided. + - second is a list of matched learnwares + """ + if anchor_flag: + return self.search_anchor_learnware(user_info) + else: + return self.search_learnware(user_info) \ No newline at end of file From d0435bc413b96961c1fd8cedb4ae203b37c77d7a Mon Sep 17 00:00:00 2001 From: Gene Date: Sat, 28 Oct 2023 22:25:07 +0800 Subject: [PATCH 24/35] [MNT] format code by black --- learnware/market/anchor/searcher.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/learnware/market/anchor/searcher.py b/learnware/market/anchor/searcher.py index 0ccdd34..cbcc7ac 100644 --- a/learnware/market/anchor/searcher.py +++ b/learnware/market/anchor/searcher.py @@ -82,7 +82,7 @@ class AnchoredSearcher(EasySearcher): - second is a list of matched learnwares """ pass - + def __call__(self, user_info: AnchoredUserInfo, anchor_flag: bool = False) -> Tuple[Any, List[Learnware]]: """Search learnwares with anchor marget - if 'anchor_flag' == True, search anchor Learnwares from anchor_learnware_list based on user_info @@ -105,4 +105,4 @@ class AnchoredSearcher(EasySearcher): if anchor_flag: return self.search_anchor_learnware(user_info) else: - return self.search_learnware(user_info) \ No newline at end of file + return self.search_learnware(user_info) From c975a1ff019480efffb10cf18c27852615ffb67c Mon Sep 17 00:00:00 2001 From: Gene Date: Sat, 28 Oct 2023 22:30:53 +0800 Subject: [PATCH 25/35] [MNT] add logger about anchor_searcher --- learnware/market/anchor/searcher.py | 3 +++ 1 file changed, 3 insertions(+) diff --git a/learnware/market/anchor/searcher.py b/learnware/market/anchor/searcher.py index cbcc7ac..ce5d489 100644 --- a/learnware/market/anchor/searcher.py +++ b/learnware/market/anchor/searcher.py @@ -2,8 +2,11 @@ from typing import List, Dict, Tuple, Any, Union from ..base import BaseUserInfo from ..easy2.searcher import EasySearcher +from ...logger import get_module_logger from ...learnware import Learnware +logger = get_module_logger("anchor_searcher") + class AnchoredUserInfo(BaseUserInfo): """ From 96b2636d9cfa84fedbd697674eb9946f0e131c34 Mon Sep 17 00:00:00 2001 From: Gene Date: Sat, 28 Oct 2023 22:33:58 +0800 Subject: [PATCH 26/35] [FIX] fix typo about evolved --- learnware/market/__init__.py | 2 +- learnware/market/evolve_anchor/__init__.py | 2 +- learnware/market/evolve_anchor/organizer.py | 2 +- 3 files changed, 3 insertions(+), 3 deletions(-) diff --git a/learnware/market/__init__.py b/learnware/market/__init__.py index bc6137e..7dd1bc6 100644 --- a/learnware/market/__init__.py +++ b/learnware/market/__init__.py @@ -1,6 +1,6 @@ from .anchor import AnchoredUserInfo, AnchoredOrganizer from .base import BaseUserInfo, LearnwareMarket, BaseChecker, BaseOrganizer, BaseSearcher -from .evolve_anchor import EvolveAnchoredOrganizer +from .evolve_anchor import EvolvedAnchoredOrganizer from .evolve import EvolvedOrganizer from .easy2 import EasyOrganizer, EasySearcher, EasySemanticChecker, EasyStatisticalChecker from .hetergeneous import HeterogeneousOrganizer, MappingFunction diff --git a/learnware/market/evolve_anchor/__init__.py b/learnware/market/evolve_anchor/__init__.py index 18624c8..a83cc4b 100644 --- a/learnware/market/evolve_anchor/__init__.py +++ b/learnware/market/evolve_anchor/__init__.py @@ -1 +1 @@ -from .organizer import EvolveAnchoredOrganizer +from .organizer import EvolvedAnchoredOrganizer diff --git a/learnware/market/evolve_anchor/organizer.py b/learnware/market/evolve_anchor/organizer.py index 04e9779..e096ec8 100644 --- a/learnware/market/evolve_anchor/organizer.py +++ b/learnware/market/evolve_anchor/organizer.py @@ -7,7 +7,7 @@ from ...logger import get_module_logger logger = get_module_logger("evolve_anchor_organizer") -class EvolveAnchoredOrganizer(AnchoredOrganizer, EvolvedOrganizer): +class EvolvedAnchoredOrganizer(AnchoredOrganizer, EvolvedOrganizer): """Organize learnwares and enable them to continuously evolve""" def __init__(self, *args, **kwargs): From 6f082ca92e075d8164ef27471707f5134117a42a Mon Sep 17 00:00:00 2001 From: Gene Date: Sat, 28 Oct 2023 23:17:17 +0800 Subject: [PATCH 27/35] [MNT] refactor EasySearcher into ExactSemantic, FuzzSemantic and Table Searcher --- learnware/market/easy2/searcher.py | 441 +++++++++++++++-------------- 1 file changed, 231 insertions(+), 210 deletions(-) diff --git a/learnware/market/easy2/searcher.py b/learnware/market/easy2/searcher.py index cd2759f..7e1297d 100644 --- a/learnware/market/easy2/searcher.py +++ b/learnware/market/easy2/searcher.py @@ -4,6 +4,7 @@ from rapidfuzz import fuzz from cvxopt import solvers, matrix from typing import Tuple, List +from .organizer import EasyOrganizer from ..base import BaseUserInfo, BaseSearcher from ...learnware import Learnware from ...specification import RKMEStatSpecification @@ -12,7 +13,182 @@ from ...logger import get_module_logger logger = get_module_logger("easy_seacher") -class EasySearcher(BaseSearcher): +class EasyExactSemanticSearcher(BaseSearcher): + def _match_semantic_spec(self, semantic_spec1, semantic_spec2): + """ + semantic_spec1: semantic spec input by user + semantic_spec2: semantic spec in database + """ + if semantic_spec1.keys() != semantic_spec2.keys(): + # sematic spec in database may contain more keys than user input + pass + + name2 = semantic_spec2["Name"]["Values"].lower() + description2 = semantic_spec2["Description"]["Values"].lower() + + for key in semantic_spec1.keys(): + v1 = semantic_spec1[key]["Values"] + v2 = semantic_spec2[key]["Values"] + + if len(v1) == 0: + # user input is empty, no need to search + continue + + if key in ("Name", "Description"): + v1 = v1.lower() + if v1 not in name2 and v1 not in description2: + return False + pass + else: + if len(v2) == 0: + # user input contains some key that is not in database + return False + + if semantic_spec1[key]["Type"] == "Class": + if isinstance(v1, list): + v1 = v1[0] + if isinstance(v2, list): + v2 = v2[0] + if v1 != v2: + return False + elif semantic_spec1[key]["Type"] == "Tag": + if not (set(v1) & set(v2)): + return False + pass + pass + pass + + return True + + def __call__(self, learnware_list: List[Learnware], user_info: BaseUserInfo) -> List[Learnware]: + match_learnwares = [] + for learnware in learnware_list: + learnware_semantic_spec = learnware.get_specification().get_semantic_spec() + user_semantic_spec = user_info.get_semantic_spec() + if self._match_semantic_spec(user_semantic_spec, learnware_semantic_spec): + match_learnwares.append(learnware) + logger.info("semantic_spec search: choose %d from %d learnwares" % (len(match_learnwares), len(learnware_list))) + return match_learnwares + + +class EasyFuzzSemanticSearcher(BaseSearcher): + def _match_semantic_spec_tag(self, semantic_spec1, semantic_spec2) -> bool: + """Judge if tags of two semantic specs are consistent + + Parameters + ---------- + semantic_spec1 : + semantic spec input by user + semantic_spec2 : + semantic spec in database + + Returns + ------- + bool + consistent (True) or not consistent (False) + """ + for key in semantic_spec1.keys(): + v1 = semantic_spec1[key]["Values"] + v2 = semantic_spec2[key]["Values"] + + if len(v1) == 0: + # user input is empty, no need to search + continue + + if key not in "Name": + if len(v2) == 0: + # user input contains some key that is not in database + return False + + if semantic_spec1[key]["Type"] == "Class": + if isinstance(v1, list): + v1 = v1[0] + if isinstance(v2, list): + v2 = v2[0] + if v1 != v2: + return False + elif semantic_spec1[key]["Type"] == "Tag": + if not (set(v1) & set(v2)): + return False + return True + + def __call__( + self, learnware_list: List[Learnware], user_info: BaseUserInfo, max_num: int = 50000, min_score: float = 75.0 + ) -> List[Learnware]: + """Search learnware by fuzzy matching of semantic spec + + Parameters + ---------- + learnware_list : List[Learnware] + The list of learnwares + user_info : BaseUserInfo + user_info contains semantic_spec + max_num : int, optional + maximum number of learnwares returned, by default 50000 + min_score : float, optional + Minimum fuzzy matching score of learnwares returned, by default 30.0 + + Returns + ------- + List[Learnware] + The list of returned learnwares + """ + matched_learnware_tag = [] + final_result = [] + user_semantic_spec = user_info.get_semantic_spec() + + for learnware in learnware_list: + learnware_semantic_spec = learnware.get_specification().get_semantic_spec() + if self._match_semantic_spec_tag(user_semantic_spec, learnware_semantic_spec): + matched_learnware_tag.append(learnware) + + if len(matched_learnware_tag) > 0: + if "Name" in user_semantic_spec: + name_user = user_semantic_spec["Name"]["Values"].lower() + if len(name_user) > 0: + # Exact search + name_list = [ + learnware.get_specification().get_semantic_spec()["Name"]["Values"].lower() + for learnware in matched_learnware_tag + ] + des_list = [ + learnware.get_specification().get_semantic_spec()["Description"]["Values"].lower() + for learnware in matched_learnware_tag + ] + + matched_learnware_exact = [] + for i in range(len(name_list)): + if name_user in name_list[i] or name_user in des_list[i]: + matched_learnware_exact.append(matched_learnware_tag[i]) + + if len(matched_learnware_exact) == 0: + # Fuzzy search + matched_learnware_fuzz, fuzz_scores = [], [] + for i in range(len(name_list)): + score_name = fuzz.partial_ratio(name_user, name_list[i]) + score_des = fuzz.partial_ratio(name_user, des_list[i]) + final_score = max(score_name, score_des) + if final_score >= min_score: + matched_learnware_fuzz.append(matched_learnware_tag[i]) + fuzz_scores.append(final_score) + + # Sort by score + sort_idx = sorted(list(range(len(fuzz_scores))), key=lambda k: fuzz_scores[k], reverse=True)[ + :max_num + ] + final_result = [matched_learnware_fuzz[idx] for idx in sort_idx] + else: + final_result = matched_learnware_exact + else: + final_result = matched_learnware_tag + else: + final_result = matched_learnware_tag + + logger.info("semantic_spec search: choose %d from %d learnwares" % (len(final_result), len(learnware_list))) + return final_result + + +class EasyTableSearcher(BaseSearcher): def _convert_dist_to_score( self, dist_list: List[float], dist_epsilon: float = 0.01, min_score: float = 0.92 ) -> List[float]: @@ -375,179 +551,60 @@ class EasySearcher(BaseSearcher): return sorted_dist_list, sorted_learnware_list - def _search_by_semantic_spec_exact( - self, learnware_list: List[Learnware], user_info: BaseUserInfo - ) -> List[Learnware]: - def match_semantic_spec(semantic_spec1, semantic_spec2): - """ - semantic_spec1: semantic spec input by user - semantic_spec2: semantic spec in database - """ - if semantic_spec1.keys() != semantic_spec2.keys(): - # sematic spec in database may contain more keys than user input - pass - - name2 = semantic_spec2["Name"]["Values"].lower() - description2 = semantic_spec2["Description"]["Values"].lower() - - for key in semantic_spec1.keys(): - v1 = semantic_spec1[key]["Values"] - v2 = semantic_spec2[key]["Values"] - - if len(v1) == 0: - # user input is empty, no need to search - continue - - if key in ("Name", "Description"): - v1 = v1.lower() - if v1 not in name2 and v1 not in description2: - return False - pass - else: - if len(v2) == 0: - # user input contains some key that is not in database - return False - - if semantic_spec1[key]["Type"] == "Class": - if isinstance(v1, list): - v1 = v1[0] - if isinstance(v2, list): - v2 = v2[0] - if v1 != v2: - return False - elif semantic_spec1[key]["Type"] == "Tag": - if not (set(v1) & set(v2)): - return False - pass - pass - pass - - return True - - match_learnwares = [] - for learnware in learnware_list: - learnware_semantic_spec = learnware.get_specification().get_semantic_spec() - user_semantic_spec = user_info.get_semantic_spec() - if match_semantic_spec(user_semantic_spec, learnware_semantic_spec): - match_learnwares.append(learnware) - logger.info("semantic_spec search: choose %d from %d learnwares" % (len(match_learnwares), len(learnware_list))) - return match_learnwares - - def _search_by_semantic_spec_fuzz( - self, learnware_list: List[Learnware], user_info: BaseUserInfo, max_num: int = 50000, min_score: float = 75.0 - ) -> List[Learnware]: - """Search learnware by fuzzy matching of semantic spec - - Parameters - ---------- - learnware_list : List[Learnware] - The list of learnwares - user_info : BaseUserInfo - user_info contains semantic_spec - max_num : int, optional - maximum number of learnwares returned, by default 50000 - min_score : float, optional - Minimum fuzzy matching score of learnwares returned, by default 30.0 - - Returns - ------- - List[Learnware] - The list of returned learnwares - """ - - def _match_semantic_spec_tag(semantic_spec1, semantic_spec2) -> bool: - """Judge if tags of two semantic specs are consistent - - Parameters - ---------- - semantic_spec1 : - semantic spec input by user - semantic_spec2 : - semantic spec in database - - Returns - ------- - bool - consistent (True) or not consistent (False) - """ - for key in semantic_spec1.keys(): - v1 = semantic_spec1[key]["Values"] - v2 = semantic_spec2[key]["Values"] - - if len(v1) == 0: - # user input is empty, no need to search - continue - - if key not in "Name": - if len(v2) == 0: - # user input contains some key that is not in database - return False - - if semantic_spec1[key]["Type"] == "Class": - if isinstance(v1, list): - v1 = v1[0] - if isinstance(v2, list): - v2 = v2[0] - if v1 != v2: - return False - elif semantic_spec1[key]["Type"] == "Tag": - if not (set(v1) & set(v2)): - return False - return True - - matched_learnware_tag = [] - final_result = [] - user_semantic_spec = user_info.get_semantic_spec() - - for learnware in learnware_list: - learnware_semantic_spec = learnware.get_specification().get_semantic_spec() - if _match_semantic_spec_tag(user_semantic_spec, learnware_semantic_spec): - matched_learnware_tag.append(learnware) + def __call__( + self, + learnware_list: List[Learnware], + user_info: BaseUserInfo, + max_search_num: int = 5, + search_method: str = "greedy", + ) -> Tuple[List[float], List[Learnware], float, List[Learnware]]: + user_rkme = user_info.stat_info["RKMEStatSpecification"] + learnware_list = self._filter_by_rkme_spec_dimension(learnware_list, user_rkme) + logger.info(f"After filter by rkme dimension, learnware_list length is {len(learnware_list)}") + + sorted_dist_list, single_learnware_list = self._search_by_rkme_spec_single(learnware_list, user_rkme) + if search_method == "auto": + mixture_dist, weight_list, mixture_learnware_list = self._search_by_rkme_spec_mixture_auto( + learnware_list, user_rkme, max_search_num + ) + elif search_method == "greedy": + mixture_dist, weight_list, mixture_learnware_list = self._search_by_rkme_spec_mixture_greedy( + learnware_list, user_rkme, max_search_num + ) + else: + logger.warning("f{search_method} not supported!") + mixture_dist = None + weight_list = [] + mixture_learnware_list = [] + + if mixture_dist is None: + sorted_score_list = self._convert_dist_to_score(sorted_dist_list) + mixture_score = None + else: + merge_score_list = self._convert_dist_to_score(sorted_dist_list + [mixture_dist]) + sorted_score_list = merge_score_list[:-1] + mixture_score = merge_score_list[-1] - if len(matched_learnware_tag) > 0: - if "Name" in user_semantic_spec: - name_user = user_semantic_spec["Name"]["Values"].lower() - if len(name_user) > 0: - # Exact search - name_list = [ - learnware.get_specification().get_semantic_spec()["Name"]["Values"].lower() - for learnware in matched_learnware_tag - ] - des_list = [ - learnware.get_specification().get_semantic_spec()["Description"]["Values"].lower() - for learnware in matched_learnware_tag - ] + logger.info(f"After search by rkme spec, learnware_list length is {len(learnware_list)}") + # filter learnware with low score + sorted_score_list, single_learnware_list = self._filter_by_rkme_spec_single( + sorted_score_list, single_learnware_list + ) - matched_learnware_exact = [] - for i in range(len(name_list)): - if name_user in name_list[i] or name_user in des_list[i]: - matched_learnware_exact.append(matched_learnware_tag[i]) + logger.info(f"After filter by rkme spec, learnware_list length is {len(learnware_list)}") + return sorted_score_list, single_learnware_list, mixture_score, mixture_learnware_list - if len(matched_learnware_exact) == 0: - # Fuzzy search - matched_learnware_fuzz, fuzz_scores = [], [] - for i in range(len(name_list)): - score_name = fuzz.partial_ratio(name_user, name_list[i]) - score_des = fuzz.partial_ratio(name_user, des_list[i]) - final_score = max(score_name, score_des) - if final_score >= min_score: - matched_learnware_fuzz.append(matched_learnware_tag[i]) - fuzz_scores.append(final_score) - # Sort by score - sort_idx = sorted(list(range(len(fuzz_scores))), key=lambda k: fuzz_scores[k], reverse=True)[ - :max_num - ] - final_result = [matched_learnware_fuzz[idx] for idx in sort_idx] - else: - final_result = matched_learnware_exact - else: - final_result = matched_learnware_tag - else: - final_result = matched_learnware_tag +class EasySearcher(BaseSearcher): + def __init__(self, organizer: EasyOrganizer = None): + super(EasySearcher, self).__init__(organizer) + self.semantic_searcher = EasyFuzzSemanticSearcher(organizer) + self.table_searcher = EasyTableSearcher(organizer) - logger.info("semantic_spec search: choose %d from %d learnwares" % (len(final_result), len(learnware_list))) - return final_result + def reset(self, organizer): + self.learnware_oganizer = organizer + self.semantic_searcher.reset(organizer) + self.table_searcher.reset(organizer) def __call__( self, user_info: BaseUserInfo, max_search_num: int = 5, search_method: str = "greedy" @@ -570,47 +627,11 @@ class EasySearcher(BaseSearcher): the fourth is the list of Learnware (mixture), the size is search_num """ learnware_list = self.learnware_oganizer.get_learnwares() - # learnware_list = self._search_by_semantic_spec_exact(learnware_list, user_info) - # if len(learnware_list) == 0: - learnware_list = self._search_by_semantic_spec_fuzz(learnware_list, user_info) + learnware_list = self.semantic_searcher(learnware_list, user_info) - if "RKMEStatSpecification" not in user_info.stat_info: - return None, learnware_list, 0.0, None - elif len(learnware_list) == 0: + if len(learnware_list) == 0: return [], [], 0.0, [] + elif "RKMEStatSpecification" in user_info.stat_info: + return self.table_searcher(learnware_list, user_info, max_search_num, search_method) else: - user_rkme = user_info.stat_info["RKMEStatSpecification"] - learnware_list = self._filter_by_rkme_spec_dimension(learnware_list, user_rkme) - logger.info(f"After filter by rkme dimension, learnware_list length is {len(learnware_list)}") - - sorted_dist_list, single_learnware_list = self._search_by_rkme_spec_single(learnware_list, user_rkme) - if search_method == "auto": - mixture_dist, weight_list, mixture_learnware_list = self._search_by_rkme_spec_mixture_auto( - learnware_list, user_rkme, max_search_num - ) - elif search_method == "greedy": - mixture_dist, weight_list, mixture_learnware_list = self._search_by_rkme_spec_mixture_greedy( - learnware_list, user_rkme, max_search_num - ) - else: - logger.warning("f{search_method} not supported!") - mixture_dist = None - weight_list = [] - mixture_learnware_list = [] - - if mixture_dist is None: - sorted_score_list = self._convert_dist_to_score(sorted_dist_list) - mixture_score = None - else: - merge_score_list = self._convert_dist_to_score(sorted_dist_list + [mixture_dist]) - sorted_score_list = merge_score_list[:-1] - mixture_score = merge_score_list[-1] - - logger.info(f"After search by rkme spec, learnware_list length is {len(learnware_list)}") - # filter learnware with low score - sorted_score_list, single_learnware_list = self._filter_by_rkme_spec_single( - sorted_score_list, single_learnware_list - ) - - logger.info(f"After filter by rkme spec, learnware_list length is {len(learnware_list)}") - return sorted_score_list, single_learnware_list, mixture_score, mixture_learnware_list + return None, learnware_list, 0.0, None From 258ab6951bfcf9ebacf05a34e9fdb61d88781ebb Mon Sep 17 00:00:00 2001 From: bxdd Date: Sun, 29 Oct 2023 13:47:43 +0800 Subject: [PATCH 28/35] [MNT] modift check_learnware --- learnware/market/base.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/learnware/market/base.py b/learnware/market/base.py index 1c26cea..338bdf5 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -101,9 +101,9 @@ class LearnwareMarket: return BaseChecker.INVALID_LEARNWARE def add_learnware( - self, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs + self, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, check_status: bool = None, **kwargs ) -> Tuple[str, bool]: - check_status = self.check_learnware(zip_path, semantic_spec, checker_names) + check_status = self.check_learnware(zip_path, semantic_spec, checker_names) if check_status is None else check_status return self.learnware_organizer.add_learnware( zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs ) @@ -115,9 +115,9 @@ class LearnwareMarket: return self.learnware_organizer.delete_learnware(id, **kwargs) def update_learnware( - self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs + self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, check_status: bool = None, **kwargs ) -> bool: - check_status = self.check_learnware(zip_path, semantic_spec, checker_names) + check_status = self.check_learnware(zip_path, semantic_spec, checker_names) if check_status is None else check_status return self.learnware_organizer.update_learnware( id, zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs ) From aac519d2cc1be7ee90d9d2e7873972ae0ff78066 Mon Sep 17 00:00:00 2001 From: bxdd Date: Sun, 29 Oct 2023 13:48:16 +0800 Subject: [PATCH 29/35] [FIX] fix bugs --- learnware/market/base.py | 7 ++----- 1 file changed, 2 insertions(+), 5 deletions(-) diff --git a/learnware/market/base.py b/learnware/market/base.py index 338bdf5..df8e501 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -1,8 +1,5 @@ -import os -import torch +import zipfile import tempfile -import traceback -import numpy as np from typing import Tuple, Any, List, Union @@ -80,7 +77,7 @@ class LearnwareMarket: z_file.extractall(tempdir) pending_learnware = get_learnware_from_dirpath( - id="pending", semantic_spec=semantic_specification, learnware_dirpath=tempdir + id="pending", semantic_spec=semantic_spec, learnware_dirpath=tempdir ) final_status = BaseChecker.INVALID_LEARNWARE From 18a06233ad6686a22283b4172c6e15f5c9d1cdf9 Mon Sep 17 00:00:00 2001 From: Gene Date: Sun, 29 Oct 2023 15:02:43 +0800 Subject: [PATCH 30/35] [MNT] modify details --- learnware/market/base.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/learnware/market/base.py b/learnware/market/base.py index df8e501..d81f1b9 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -98,9 +98,9 @@ class LearnwareMarket: return BaseChecker.INVALID_LEARNWARE def add_learnware( - self, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, check_status: bool = None, **kwargs + self, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs ) -> Tuple[str, bool]: - check_status = self.check_learnware(zip_path, semantic_spec, checker_names) if check_status is None else check_status + check_status = self.check_learnware(zip_path, semantic_spec, checker_names) return self.learnware_organizer.add_learnware( zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs ) @@ -112,9 +112,9 @@ class LearnwareMarket: return self.learnware_organizer.delete_learnware(id, **kwargs) def update_learnware( - self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, check_status: bool = None, **kwargs + self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs ) -> bool: - check_status = self.check_learnware(zip_path, semantic_spec, checker_names) if check_status is None else check_status + check_status = self.check_learnware(zip_path, semantic_spec, checker_names) return self.learnware_organizer.update_learnware( id, zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs ) From c10de4334254728477fbd51c1ef24c357f0e46e9 Mon Sep 17 00:00:00 2001 From: Gene Date: Sun, 29 Oct 2023 17:56:22 +0800 Subject: [PATCH 31/35] [FIX] fix bugs in SemanticSearcher --- learnware/market/easy2/searcher.py | 8 ++++---- 1 file changed, 4 insertions(+), 4 deletions(-) diff --git a/learnware/market/easy2/searcher.py b/learnware/market/easy2/searcher.py index 7e1297d..9d758fc 100644 --- a/learnware/market/easy2/searcher.py +++ b/learnware/market/easy2/searcher.py @@ -27,8 +27,8 @@ class EasyExactSemanticSearcher(BaseSearcher): description2 = semantic_spec2["Description"]["Values"].lower() for key in semantic_spec1.keys(): - v1 = semantic_spec1[key]["Values"] - v2 = semantic_spec2[key]["Values"] + v1 = semantic_spec1[key].get("Values", "") + v2 = semantic_spec2[key].get("Values", "") if len(v1) == 0: # user input is empty, no need to search @@ -88,8 +88,8 @@ class EasyFuzzSemanticSearcher(BaseSearcher): consistent (True) or not consistent (False) """ for key in semantic_spec1.keys(): - v1 = semantic_spec1[key]["Values"] - v2 = semantic_spec2[key]["Values"] + v1 = semantic_spec1[key].get("Values", "") + v2 = semantic_spec2[key].get("Values", "") if len(v1) == 0: # user input is empty, no need to search From 715049c19b1ffbc8640ec8b875bd21d18d4a4493 Mon Sep 17 00:00:00 2001 From: Gene Date: Sun, 29 Oct 2023 17:57:10 +0800 Subject: [PATCH 32/35] [MNT] modify add_learnware and update_learnware --- learnware/market/base.py | 47 ++++++++++++++-- learnware/market/easy2/organizer.py | 84 ++++++++++++++--------------- 2 files changed, 84 insertions(+), 47 deletions(-) diff --git a/learnware/market/base.py b/learnware/market/base.py index d81f1b9..4b1332c 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -99,7 +99,24 @@ class LearnwareMarket: def add_learnware( self, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs - ) -> Tuple[str, bool]: + ) -> Tuple[str, int]: + """Add a learnware into the market. + + Parameters + ---------- + zip_path : str + Filepath for learnware model, a zipped file. + semantic_spec : dict + semantic_spec for new learnware, in dictionary format. + checker_names : List[str], optional + List contains checker names, by default None + + Returns + ------- + Tuple[str, int] + - str indicating model_id + - int indicating the final learnware check_status + """ check_status = self.check_learnware(zip_path, semantic_spec, checker_names) return self.learnware_organizer.add_learnware( zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs @@ -112,9 +129,31 @@ class LearnwareMarket: return self.learnware_organizer.delete_learnware(id, **kwargs) def update_learnware( - self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, **kwargs - ) -> bool: - check_status = self.check_learnware(zip_path, semantic_spec, checker_names) + self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, check_status: int = None, **kwargs + ) -> int: + """Update learnware with zip_path and semantic_specification + + Parameters + ---------- + id : str + Learnware id + zip_path : str + Filepath for learnware model, a zipped file. + semantic_spec : dict + semantic_spec for new learnware, in dictionary format. + checker_names : List[str], optional + List contains checker names, by default None. + check_status : int, optional + A flag indicating whether the learnware is usable, by default None. + + Returns + ------- + int + The final learnware check_status. + """ + update_status = self.check_learnware(zip_path, semantic_spec, checker_names) + check_status = update_status if check_status is None or update_status == BaseChecker.INVALID_LEARNWARE else check_status + return self.learnware_organizer.update_learnware( id, zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs ) diff --git a/learnware/market/easy2/organizer.py b/learnware/market/easy2/organizer.py index 55780e3..3a78794 100644 --- a/learnware/market/easy2/organizer.py +++ b/learnware/market/easy2/organizer.py @@ -37,7 +37,6 @@ class EasyOrganizer(BaseOrganizer): bool A flag indicating whether the market is reload successfully. """ - self.market_store_path = os.path.join(conf.market_root_path, self.market_id) self.learnware_pool_path = os.path.join(self.market_store_path, "learnware_pool") self.learnware_zip_pool_path = os.path.join(self.learnware_pool_path, "zips") @@ -70,33 +69,33 @@ class EasyOrganizer(BaseOrganizer): ) = self.dbops.load_market() def add_learnware( - self, zip_path: str, semantic_spec: dict, id: str = None, check_status: int = None - ) -> Tuple[str, bool]: + self, zip_path: str, semantic_spec: dict, check_status: int + ) -> Tuple[str, int]: """Add a learnware into the market. - .. note:: - - Given a prediction of a certain time, all signals before this time will be prepared well. - - Parameters ---------- zip_path : str Filepath for learnware model, a zipped file. semantic_spec : dict semantic_spec for new learnware, in dictionary format. + check_status: int + A flag indicating whether the learnware is usable. Returns ------- Tuple[str, int] - str indicating model_id - - int indicating what the flag of learnware is added. - + - int indicating the final learnware check_status """ + if check_status == BaseChecker.INVALID_LEARNWARE: + logger.warning("Learnware is invalid!") + return None, BaseChecker.INVALID_LEARNWARE + semantic_spec = copy.deepcopy(semantic_spec) logger.info("Get new learnware from %s" % (zip_path)) - id = id if id is not None else "%08d" % (self.count) + id = "%08d" % (self.count) target_zip_dir = os.path.join(self.learnware_zip_pool_path, "%s.zip" % (id)) target_folder_dir = os.path.join(self.learnware_folder_pool_path, id) copyfile(zip_path, target_zip_dir) @@ -168,44 +167,43 @@ class EasyOrganizer(BaseOrganizer): return True def update_learnware(self, id: str, zip_path: str = None, semantic_spec: dict = None, check_status: int = None): - """update learnware with zip_path and semantic_specification - TODO: update should pass the semantic check too + """Update learnware with zip_path, semantic_specification and check_status Parameters ---------- id : str - _description_ + Learnware id zip_path : str, optional - _description_, by default None + Filepath for learnware model, a zipped file. semantic_spec : dict, optional - _description_, by default None + semantic_spec for new learnware, in dictionary format. check_status : int, optional - _description_, by default None + A flag indicating whether the learnware is usable. Returns ------- - _type_ - _description_ + int + The final learnware check_status. """ - assert ( - zip_path is None and semantic_spec is None - ), f"at least one of 'zip_path' and 'semantic_spec' should not be None when update learnware" - assert check_status != BaseChecker.INVALID_LEARNWARE, f"'check_status' can not be INVALID_LEARNWARE" - - if zip_path is None and check_status is not None: - logger.warning("check_status will be ignored when zip_path is None for learnware update") - + if check_status == BaseChecker.INVALID_LEARNWARE: + logger.warning("Learnware is invalid!") + return BaseChecker.INVALID_LEARNWARE + + if zip_path is None and semantic_spec is None and check_status is None: + logger.warning("At least one of 'zip_path', 'semantic_spec' and 'check_status' should not be None when update learnware") + return BaseChecker.INVALID_LEARNWARE + + # Update semantic_specification learnware_zippath = self.learnware_zip_list[id] if zip_path is None else zip_path semantic_spec = ( self.learnware_list[id].get_specification().get_semantic_spec() if semantic_spec is None else semantic_spec ) - self.dbops.update_learnware_semantic_specification(id, semantic_spec) - + + # Update zip path target_zip_dir = self.learnware_zip_list[id] target_folder_dir = self.learnware_folder_list[id] - - if check_status is None and zip_path is not None: + if zip_path is not None: with tempfile.TemporaryDirectory(prefix="learnware_") as tempdir: with zipfile.ZipFile(zip_path, "r") as z_file: z_file.extractall(tempdir) @@ -219,21 +217,21 @@ class EasyOrganizer(BaseOrganizer): if new_learnware is None: return BaseChecker.INVALID_LEARNWARE - - learnwere_status = BaseChecker.NONUSABLE_LEARNWARE - else: - learnwere_status = self.use_flags[id] if zip_path is None else check_status - - copyfile(zip_path, target_zip_dir) - with zipfile.ZipFile(target_zip_dir, "r") as z_file: - z_file.extractall(target_folder_dir) - + + copyfile(zip_path, target_zip_dir) + with zipfile.ZipFile(target_zip_dir, "r") as z_file: + z_file.extractall(target_folder_dir) + + # Update check_status + self.use_flags[id] = self.use_flags[id] if check_status is None else check_status + self.dbops.update_learnware_use_flag(id, self.use_flags[id]) + + # Update learnware list self.learnware_list[id] = get_learnware_from_dirpath( id=id, semantic_spec=semantic_spec, learnware_dirpath=target_folder_dir ) - self.use_flags[id] = learnwere_status - self.dbops.update_learnware_use_flag(id, learnwere_status) - return learnwere_status + + return self.use_flags[id] def get_learnware_by_ids(self, ids: Union[str, List[str]]) -> Union[Learnware, List[Learnware]]: """Search learnware by id or list of ids. From e1093f3beb5ee7488cd7ce4bc282c4eaaf77ff73 Mon Sep 17 00:00:00 2001 From: Gene Date: Sun, 29 Oct 2023 17:58:10 +0800 Subject: [PATCH 33/35] [MNT] expand testing scope --- tests/test_market/test_easy.py | 18 ++++++++++++++++-- 1 file changed, 16 insertions(+), 2 deletions(-) diff --git a/tests/test_market/test_easy.py b/tests/test_market/test_easy.py index 2aa674c..aa7cad8 100644 --- a/tests/test_market/test_easy.py +++ b/tests/test_market/test_easy.py @@ -26,6 +26,12 @@ user_semantic = { "Scenario": {"Values": ["Education"], "Type": "Tag"}, "Description": {"Values": "", "Type": "String"}, "Name": {"Values": "", "Type": "String"}, + "Output": { + "Dimension": 10, + "Description": { + "0": "the probability of the label is zero", + }, + } } @@ -85,28 +91,36 @@ class TestMarket(unittest.TestCase): self.zip_path_list.append(zip_file) - def test_upload_delete_learnware(self, learnware_num=5, delete=False): + def test_upload_delete_learnware(self, learnware_num=5, delete=True): easy_market = self._init_learnware_market() self.test_prepare_learnware_randomly(learnware_num) + self.learnware_num = learnware_num print("Total Item:", len(easy_market)) + assert len(easy_market) == 0, f"The market should be empty!" for idx, zip_path in enumerate(self.zip_path_list): semantic_spec = copy.deepcopy(user_semantic) semantic_spec["Name"]["Values"] = "learnware_%d" % (idx) semantic_spec["Description"]["Values"] = "test_learnware_number_%d" % (idx) - semantic_spec["Output"] = {"Dimension": 1, "Description": {"0": "The label of the hand-written digit."}} easy_market.add_learnware(zip_path, semantic_spec) print("Total Item:", len(easy_market)) + assert len(easy_market) == self.learnware_num, f"The number of learnwares must be {self.learnware_num}!" + curr_inds = easy_market.get_learnware_ids() print("Available ids After Uploading Learnwares:", curr_inds) + assert len(curr_inds) == self.learnware_num, f"The number of learnwares must be {self.learnware_num}!" if delete: for learnware_id in curr_inds: easy_market.delete_learnware(learnware_id) + self.learnware_num -= 1 + assert len(easy_market) == self.learnware_num, f"The number of learnwares must be {self.learnware_num}!" + curr_inds = easy_market.get_learnware_ids() print("Available ids After Deleting Learnwares:", curr_inds) + assert len(curr_inds) == 0, f"The market should be empty!" return easy_market From ad48d6d3dc6ed208b2d40aaf01974becb8f96396 Mon Sep 17 00:00:00 2001 From: Gene Date: Sun, 29 Oct 2023 19:50:56 +0800 Subject: [PATCH 34/35] [MNT] expand testing scope --- tests/test_market/test_easy.py | 29 ++++++++++++++++------------- 1 file changed, 16 insertions(+), 13 deletions(-) diff --git a/tests/test_market/test_easy.py b/tests/test_market/test_easy.py index aa7cad8..80268c9 100644 --- a/tests/test_market/test_easy.py +++ b/tests/test_market/test_easy.py @@ -127,30 +127,32 @@ class TestMarket(unittest.TestCase): def test_search_semantics(self, learnware_num=5): easy_market = self.test_upload_delete_learnware(learnware_num, delete=False) print("Total Item:", len(easy_market)) - - test_folder = os.path.join(curr_root, "test_semantics") - - # unzip -o -q zip_path -d unzip_dir - if os.path.exists(test_folder): - rmtree(test_folder) - os.makedirs(test_folder, exist_ok=True) - - with zipfile.ZipFile(self.zip_path_list[0], "r") as zip_obj: - zip_obj.extractall(path=test_folder) + assert len(easy_market) == self.learnware_num, f"The number of learnwares must be {self.learnware_num}!" semantic_spec = copy.deepcopy(user_semantic) semantic_spec["Name"]["Values"] = f"learnware_{learnware_num - 1}" - semantic_spec["Description"]["Values"] = f"test_learnware_number_{learnware_num - 1}" user_info = BaseUserInfo(semantic_spec=semantic_spec) _, single_learnware_list, _, _ = easy_market.search_learnware(user_info) print("User info:", user_info.get_semantic_spec()) print(f"Search result:") + assert len(single_learnware_list) == 1, f"Exact semantic search failed!" for learnware in single_learnware_list: - print("Choose learnware:", learnware.id, learnware.get_specification().get_semantic_spec()) + semantic_spec1 = learnware.get_specification().get_semantic_spec() + print("Choose learnware:", learnware.id, semantic_spec1) + assert semantic_spec1["Name"]["Values"] == semantic_spec["Name"]["Values"], f"Exact semantic search failed!" + + semantic_spec["Name"]["Values"] = "laernwaer" + user_info = BaseUserInfo(semantic_spec=semantic_spec) + _, single_learnware_list, _, _ = easy_market.search_learnware(user_info) - rmtree(test_folder) # rm -r test_folder + print("User info:", user_info.get_semantic_spec()) + print(f"Search result:") + assert len(single_learnware_list) == self.learnware_num, f"Fuzzy semantic search failed!" + for learnware in single_learnware_list: + semantic_spec1 = learnware.get_specification().get_semantic_spec() + print("Choose learnware:", learnware.id, semantic_spec1) def test_stat_search(self, learnware_num=5): easy_market = self.test_upload_delete_learnware(learnware_num, delete=False) @@ -178,6 +180,7 @@ class TestMarket(unittest.TestCase): mixture_learnware_list, ) = easy_market.search_learnware(user_info) + assert len(single_learnware_list) == self.learnware_num, f"Statistical search failed!" print(f"search result of user{idx}:") for score, learnware in zip(sorted_score_list, single_learnware_list): print(f"score: {score}, learnware_id: {learnware.id}") From 898b4d5c97862cab1fa68a8de3fd8e7917914fba Mon Sep 17 00:00:00 2001 From: Gene Date: Sun, 29 Oct 2023 19:51:14 +0800 Subject: [PATCH 35/35] [MNT] format code by black --- learnware/market/base.py | 12 ++++++++++-- learnware/market/easy2/organizer.py | 20 ++++++++++---------- tests/test_market/test_easy.py | 8 ++++---- 3 files changed, 24 insertions(+), 16 deletions(-) diff --git a/learnware/market/base.py b/learnware/market/base.py index 4b1332c..927d4fe 100644 --- a/learnware/market/base.py +++ b/learnware/market/base.py @@ -129,7 +129,13 @@ class LearnwareMarket: return self.learnware_organizer.delete_learnware(id, **kwargs) def update_learnware( - self, id: str, zip_path: str, semantic_spec: dict, checker_names: List[str] = None, check_status: int = None, **kwargs + self, + id: str, + zip_path: str, + semantic_spec: dict, + checker_names: List[str] = None, + check_status: int = None, + **kwargs, ) -> int: """Update learnware with zip_path and semantic_specification @@ -152,7 +158,9 @@ class LearnwareMarket: The final learnware check_status. """ update_status = self.check_learnware(zip_path, semantic_spec, checker_names) - check_status = update_status if check_status is None or update_status == BaseChecker.INVALID_LEARNWARE else check_status + check_status = ( + update_status if check_status is None or update_status == BaseChecker.INVALID_LEARNWARE else check_status + ) return self.learnware_organizer.update_learnware( id, zip_path=zip_path, semantic_spec=semantic_spec, check_status=check_status, **kwargs diff --git a/learnware/market/easy2/organizer.py b/learnware/market/easy2/organizer.py index 3a78794..830b5d3 100644 --- a/learnware/market/easy2/organizer.py +++ b/learnware/market/easy2/organizer.py @@ -68,9 +68,7 @@ class EasyOrganizer(BaseOrganizer): self.count, ) = self.dbops.load_market() - def add_learnware( - self, zip_path: str, semantic_spec: dict, check_status: int - ) -> Tuple[str, int]: + def add_learnware(self, zip_path: str, semantic_spec: dict, check_status: int) -> Tuple[str, int]: """Add a learnware into the market. Parameters @@ -91,7 +89,7 @@ class EasyOrganizer(BaseOrganizer): if check_status == BaseChecker.INVALID_LEARNWARE: logger.warning("Learnware is invalid!") return None, BaseChecker.INVALID_LEARNWARE - + semantic_spec = copy.deepcopy(semantic_spec) logger.info("Get new learnware from %s" % (zip_path)) @@ -188,9 +186,11 @@ class EasyOrganizer(BaseOrganizer): if check_status == BaseChecker.INVALID_LEARNWARE: logger.warning("Learnware is invalid!") return BaseChecker.INVALID_LEARNWARE - + if zip_path is None and semantic_spec is None and check_status is None: - logger.warning("At least one of 'zip_path', 'semantic_spec' and 'check_status' should not be None when update learnware") + logger.warning( + "At least one of 'zip_path', 'semantic_spec' and 'check_status' should not be None when update learnware" + ) return BaseChecker.INVALID_LEARNWARE # Update semantic_specification @@ -199,7 +199,7 @@ class EasyOrganizer(BaseOrganizer): self.learnware_list[id].get_specification().get_semantic_spec() if semantic_spec is None else semantic_spec ) self.dbops.update_learnware_semantic_specification(id, semantic_spec) - + # Update zip path target_zip_dir = self.learnware_zip_list[id] target_folder_dir = self.learnware_folder_list[id] @@ -217,15 +217,15 @@ class EasyOrganizer(BaseOrganizer): if new_learnware is None: return BaseChecker.INVALID_LEARNWARE - + copyfile(zip_path, target_zip_dir) with zipfile.ZipFile(target_zip_dir, "r") as z_file: z_file.extractall(target_folder_dir) - + # Update check_status self.use_flags[id] = self.use_flags[id] if check_status is None else check_status self.dbops.update_learnware_use_flag(id, self.use_flags[id]) - + # Update learnware list self.learnware_list[id] = get_learnware_from_dirpath( id=id, semantic_spec=semantic_spec, learnware_dirpath=target_folder_dir diff --git a/tests/test_market/test_easy.py b/tests/test_market/test_easy.py index 80268c9..5f22729 100644 --- a/tests/test_market/test_easy.py +++ b/tests/test_market/test_easy.py @@ -31,7 +31,7 @@ user_semantic = { "Description": { "0": "the probability of the label is zero", }, - } + }, } @@ -117,7 +117,7 @@ class TestMarket(unittest.TestCase): easy_market.delete_learnware(learnware_id) self.learnware_num -= 1 assert len(easy_market) == self.learnware_num, f"The number of learnwares must be {self.learnware_num}!" - + curr_inds = easy_market.get_learnware_ids() print("Available ids After Deleting Learnwares:", curr_inds) assert len(curr_inds) == 0, f"The market should be empty!" @@ -140,9 +140,9 @@ class TestMarket(unittest.TestCase): assert len(single_learnware_list) == 1, f"Exact semantic search failed!" for learnware in single_learnware_list: semantic_spec1 = learnware.get_specification().get_semantic_spec() - print("Choose learnware:", learnware.id, semantic_spec1) + print("Choose learnware:", learnware.id, semantic_spec1) assert semantic_spec1["Name"]["Values"] == semantic_spec["Name"]["Values"], f"Exact semantic search failed!" - + semantic_spec["Name"]["Values"] = "laernwaer" user_info = BaseUserInfo(semantic_spec=semantic_spec) _, single_learnware_list, _, _ = easy_market.search_learnware(user_info)