diff --git a/.gitignore b/.gitignore index d22ea69..39ba56e 100644 --- a/.gitignore +++ b/.gitignore @@ -43,5 +43,4 @@ cache/ tmp/ learnware_pool/ PFS/ -data/ -learnware/market/hetergeneous/.learnware/* \ No newline at end of file +data/ \ No newline at end of file diff --git a/learnware/market/__init__.py b/learnware/market/__init__.py index be3e1ac..b850f5c 100644 --- a/learnware/market/__init__.py +++ b/learnware/market/__init__.py @@ -3,7 +3,7 @@ from .base import BaseUserInfo, LearnwareMarket, BaseChecker, BaseOrganizer, Bas from .evolve_anchor import EvolvedAnchoredOrganizer from .evolve import EvolvedOrganizer from .easy import EasyOrganizer, EasySearcher, EasySemanticChecker, EasyStatChecker -from .hetergeneous import HeteroMapTableOrganizer, HeteroSearcher +from .heterogeneous import HeteroMapTableOrganizer, HeteroSearcher from .classes import CondaChecker from .module import instantiate_learnware_market diff --git a/learnware/market/hetergeneous/database_ops.py b/learnware/market/hetergeneous/database_ops.py deleted file mode 100644 index 5d8461a..0000000 --- a/learnware/market/hetergeneous/database_ops.py +++ /dev/null @@ -1,177 +0,0 @@ -import json -import os - -from learnware.learnware import get_learnware_from_dirpath -from learnware.logger import get_module_logger -from sqlalchemy import (Column, DateTime, Integer, String, Text, create_engine, - text) -from sqlalchemy.ext.declarative import declarative_base - -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 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 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;")) - - learnware_list = {} - zip_list = {} - folder_list = {} - use_flags = {} - 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 - use_flags[id] = use_flag - max_count = max(max_count, int(id)) - pass - - return learnware_list, zip_list, folder_list, use_flags, max_count + 1 - pass - - pass diff --git a/learnware/market/hetergeneous/__init__.py b/learnware/market/heterogeneous/__init__.py similarity index 100% rename from learnware/market/hetergeneous/__init__.py rename to learnware/market/heterogeneous/__init__.py diff --git a/learnware/market/hetergeneous/organizer.py b/learnware/market/heterogeneous/organizer.py similarity index 100% rename from learnware/market/hetergeneous/organizer.py rename to learnware/market/heterogeneous/organizer.py diff --git a/learnware/market/hetergeneous/organizer/__init__.py b/learnware/market/heterogeneous/organizer/__init__.py similarity index 74% rename from learnware/market/hetergeneous/organizer/__init__.py rename to learnware/market/heterogeneous/organizer/__init__.py index c202ccd..9933c19 100644 --- a/learnware/market/hetergeneous/organizer/__init__.py +++ b/learnware/market/heterogeneous/organizer/__init__.py @@ -17,9 +17,9 @@ from ....learnware import Learnware, get_learnware_from_dirpath from ....logger import get_module_logger from ....specification.system import HeteroSpecification from ...base import BaseChecker, BaseUserInfo -from ...easy2 import EasyOrganizer -from ..database_ops import DatabaseOperations -from .config import C as conf +from ...easy import EasyOrganizer +from ...easy.database_ops import DatabaseOperations +from ....config import C as conf from .hetero_mapping import HeteroMapping, Trainer logger = get_module_logger("hetero_market") @@ -27,12 +27,12 @@ logger = get_module_logger("hetero_market") class HeteroMapTableOrganizer(EasyOrganizer): def reload_market(self, rebuild=False, auto_update_limit=100): - self.market_store_path = os.path.join(conf.hetero_root_path, self.market_id) - self.market_mapping_path = os.path.join(self.market_store_path, conf.market_model_path) + self.market_store_path = os.path.join(conf.root_path, self.market_id) + self.market_mapping_path = os.path.join(self.market_store_path, "model.bin") self.learnware_pool_path = os.path.join(self.market_store_path, "learnware_pool") self.learnware_zip_pool_path = os.path.join(self.market_store_path, "zips") self.learnware_folder_pool_path = os.path.join(self.market_store_path, "unzipped_learnwares") - self.hetero_mappings_path = os.path.join(self.market_store_path, conf.heter_mapping_path) + self.hetero_mappings_path = os.path.join(self.market_store_path, "hetero_mappings") self.learnware_list = {} # id:learnware self.learnware_zip_list = {} self.learnware_folder_list = {} @@ -41,8 +41,6 @@ class HeteroMapTableOrganizer(EasyOrganizer): self.dbops = DatabaseOperations(conf.database_url, "market_" + self.market_id) self.auto_update = False self.auto_update_limit = auto_update_limit - self.auto_update_lock = mp.Lock() - self.is_training_in_progress = mp.Value('i', 0) if rebuild: logger.warning("Warning! You are trying to clear current database!") @@ -75,7 +73,6 @@ class HeteroMapTableOrganizer(EasyOrganizer): self.market_mapping = HeteroMapping() def reset(self, market_id=None, auto_update=False, auto_update_limit=None, **kwargs): - # model training arguments(model architecture + optimization) set via self.reset self.auto_update = auto_update self.market_id = market_id self.training_args = kwargs @@ -126,42 +123,45 @@ class HeteroMapTableOrganizer(EasyOrganizer): self.learnware_zip_list[learnware_id] = target_zip_dir self.learnware_folder_list[learnware_id] = target_folder_dir self.use_flags[learnware_id] = learnwere_status - self.count += 1 - - with self.auto_update_lock: - if self.auto_update and not self.is_training_in_progress.value and self.count - self.last_trained_learnware_num >= self.auto_update_limit: - self.is_training_in_progress.value = 1 - curr_learnware_list = copy.deepcopy(self.learnware_list) - train_process = mp.Process(target=self.train, args=(curr_learnware_list.values(),)) - train_process.start() - # train_process.join() + self.count += 1 + + if self.auto_update and self.count - self.last_trained_learnware_num == self.auto_update_limit + 1: + logger.warning(f"Leanwares for training: {self.get_learnware_ids()}") + + updated_market_mapping = self.train( + learnware_list=self.learnware_list.values(), + save_dir=self.market_store_path, + **self.training_args + ) + + logger.warning(f"Market mapping train completed. Now update HeteroSpecification for {self.get_learnware_ids()}") + + self.market_mapping = updated_market_mapping + self._update_learnware_list(self.learnware_list.values()) + self.last_trained_learnware_num = self.count return learnware_id, learnwere_status - def train(self, learnware_list: List[Learnware] = None): - learnware_list = learnware_list or self.learnware_list.values() - logger.warning(f"Leanwares for training: {[learnware.id for learnware in learnware_list]}") - allset = self._learnwares_to_dataframes(learnware_list) - self.market_mapping = HeteroMapping(**self.training_args) + @staticmethod + def train(learnware_list: List[Learnware] = None, save_dir: str = None, **kwargs): + allset = HeteroMapTableOrganizer._learnwares_to_dataframes(learnware_list) + market_mapping = HeteroMapping(**kwargs) market_mapping_trainer = Trainer( - model=self.market_mapping, + model=market_mapping, train_set_list=allset, - collate_fn=self.market_mapping.collate_fn, - **self.training_args, + collate_fn=market_mapping.collate_fn, + **kwargs, ) - market_mapping_trainer.train() - # auto save whenever market model retrained - market_mapping_trainer.save_model(output_dir=self.market_store_path) - - # essential hetero-mapping update for each market learnware when market model retrained - self._update_learnware_list(learnware_list) - self.last_trained_learnware_num = self.count + market_mapping_trainer.train() + market_mapping_trainer.save_model(output_dir=save_dir) - logger.warning(f"Updataed Specification For: {[learnware.id for learnware in learnware_list]}") + return market_mapping - with self.auto_update_lock: - self.is_training_in_progress.value = 0 + ############################################ + # save_model & generateing new specification + # should be moved out of train thread + ############################################ def _update_learnware_list(self, learnware_list: List[Learnware]): try: @@ -169,7 +169,7 @@ class HeteroMapTableOrganizer(EasyOrganizer): hetero_spec_path = os.path.join(self.hetero_mappings_path, f"{learnware.id}.npy") self._update_learnware_specification(learnware, save_path=hetero_spec_path) except Exception as err: - logger.warning(f"Update learnware HeteroSpecification failed! Due to {err}") + logger.warning(f"Update HeteroSpecification failed! Due to {err}") def _update_learnware_specification(self, learnware: Learnware, save_path: str) -> Learnware: specification = learnware.specification @@ -178,16 +178,16 @@ class HeteroMapTableOrganizer(EasyOrganizer): learnware_hetero_spec = self.market_mapping.hetero_mapping(learnware_rkme, learnware_features) learnware.update_stat_spec("HeteroSpecification", learnware_hetero_spec) - # custom hetero spec save path? learnware_hetero_spec.save(save_path) def generate_hetero_map_spec(self, user_info: BaseUserInfo) -> HeteroSpecification: user_rkme = user_info.stat_info["RKMETableSpecification"] - user_features = user_info.semantic_spec["Input"]["Description"].values() + user_features = user_info.get_semantic_spec()["Input"]["Description"].values() user_hetero_spec = self.market_mapping.hetero_mapping(user_rkme, user_features) return user_hetero_spec - def _learnwares_to_dataframes(self, learnware_list: List[Learnware]) -> List[pd.DataFrame]: + @staticmethod + def _learnwares_to_dataframes(learnware_list: List[Learnware]) -> List[pd.DataFrame]: learnware_df_dict = defaultdict(list) for learnware in learnware_list: specification = learnware.get_specification() @@ -198,7 +198,4 @@ class HeteroMapTableOrganizer(EasyOrganizer): learnware_df_dict[tuple(sorted(learnware_features))].append(learnware_df) merged_dfs = [pd.concat(dfs) for dfs in learnware_df_dict.values()] - return merged_dfs - - def save(self, save_path): - return NotImplementedError \ No newline at end of file + return merged_dfs \ No newline at end of file diff --git a/learnware/market/hetergeneous/organizer/hetero_mapping/__init__.py b/learnware/market/heterogeneous/organizer/hetero_mapping/__init__.py similarity index 100% rename from learnware/market/hetergeneous/organizer/hetero_mapping/__init__.py rename to learnware/market/heterogeneous/organizer/hetero_mapping/__init__.py diff --git a/learnware/market/hetergeneous/organizer/hetero_mapping/feature_extractor.py b/learnware/market/heterogeneous/organizer/hetero_mapping/feature_extractor.py similarity index 100% rename from learnware/market/hetergeneous/organizer/hetero_mapping/feature_extractor.py rename to learnware/market/heterogeneous/organizer/hetero_mapping/feature_extractor.py diff --git a/learnware/market/hetergeneous/organizer/hetero_mapping/trainer.py b/learnware/market/heterogeneous/organizer/hetero_mapping/trainer.py similarity index 100% rename from learnware/market/hetergeneous/organizer/hetero_mapping/trainer.py rename to learnware/market/heterogeneous/organizer/hetero_mapping/trainer.py diff --git a/learnware/market/hetergeneous/searcher.py b/learnware/market/heterogeneous/searcher.py similarity index 80% rename from learnware/market/hetergeneous/searcher.py rename to learnware/market/heterogeneous/searcher.py index 9d7e3a7..4ed3444 100644 --- a/learnware/market/hetergeneous/searcher.py +++ b/learnware/market/heterogeneous/searcher.py @@ -1,10 +1,12 @@ from typing import Tuple, List, Union +import numpy as np + from ...learnware import Learnware from ...logger import get_module_logger from ...specification import HeteroSpecification from ..base import BaseSearcher, BaseUserInfo -from ..easy2 import EasySearcher +from ..easy import EasySearcher from ..utils import parse_specification_type from .organizer import HeteroMapTableOrganizer @@ -38,7 +40,7 @@ class HeteroMapTableSearcher(EasySearcher): ) -> Tuple[List[float], List[Learnware]]: hetero_spec_list = [learnware.specification.get_stat_spec_by_name("HeteroSpecification") for learnware in learnware_list] mmd_dist_list = [] - for hetero_spec in hetero_spec_list: + for idx, hetero_spec in enumerate(hetero_spec_list): mmd_dist = hetero_spec.dist(user_hetero_spec) mmd_dist_list.append(mmd_dist) @@ -83,15 +85,6 @@ class HeteroMapTableSearcher(EasySearcher): logger.info(f"After filter by hetero spec, learnware_list length is {len(single_learnware_list)}") return sorted_score_list, single_learnware_list, None, None - - # for learnware in learnware_list: - # learnware_hetero_spec = learnware.specification.get_stat_spec_by_name("HeteroSpecification") - # mmd_dist = learnware_hetero_spec.dist(user_hetero_spec) - # if target_learnware is None or mmd_dist < min_dist: - # min_dist = mmd_dist - # target_learnware = learnware - # return target_learnware - def reset(self, organizer): self.learnware_oganizer = organizer @@ -103,6 +96,26 @@ class HeteroSearcher(EasySearcher): def reset(self, organizer): super().reset(organizer) self.hetero_stat_searcher.reset(organizer) + + @staticmethod + def check_user_info(user_info: BaseUserInfo): + try: + user_stat_spec = user_info.get_stat_info("RKMETableSpecification") + user_input_shape = user_stat_spec.get_z().shape[1] + + user_input_description = user_info.get_semantic_spec()["Input"] + + user_description_dim = int(user_input_description["Dimension"]) + user_description_feature_num = len(user_input_description["Description"]) + + if user_input_shape != user_description_dim or user_input_shape != user_description_feature_num: + logger.warning("User data feature dimensions mismatch with semantic specification") + return False + + return True + except: + logger.info(f"No heterogeneous search information provided. Use homogeneous search instead.") + return False def __call__( self, user_info: BaseUserInfo, check_status: int = None, max_search_num: int = 5, search_method: str = "greedy" @@ -114,7 +127,7 @@ class HeteroSearcher(EasySearcher): return [], [], 0.0, [] if parse_specification_type(stat_specs=user_info.stat_info) is not None: - if "Input" in user_info.semantic_spec and user_info.semantic_spec["Input"]["Description"] is not None: + if self.check_user_info(user_info): return self.hetero_stat_searcher(learnware_list, user_info) else: return self.stat_searcher(learnware_list, user_info, max_search_num, search_method) diff --git a/learnware/market/module.py b/learnware/market/module.py index f0903e4..945bbaf 100644 --- a/learnware/market/module.py +++ b/learnware/market/module.py @@ -1,6 +1,6 @@ from .base import LearnwareMarket from .easy import EasyOrganizer, EasySearcher, EasySemanticChecker, EasyStatChecker -from .hetergeneous import HeteroMapTableOrganizer, HeteroSearcher +from .heterogeneous import HeteroMapTableOrganizer, HeteroSearcher MARKET_CONFIG = { "easy": { diff --git a/learnware/specification/system/heter_table.py b/learnware/specification/system/heter_table.py index f721d8f..a574daf 100644 --- a/learnware/specification/system/heter_table.py +++ b/learnware/specification/system/heter_table.py @@ -29,10 +29,10 @@ class HeteroSpecification(SystemStatsSpecification): super(HeteroSpecification, self).__init__(type=self.__class__.__name__) def get_z(self) -> np.ndarray: - return self.z.detach().cpu().numpy + return self.z.detach().cpu().numpy() def get_beta(self) -> np.ndarray: - return self.beta.detach().cpu().numpy + return self.beta.detach().cpu().numpy() def generate_stat_spec_from_system(self, heter_embedding: np.ndarray, rkme_spec: RKMETableSpecification): self.beta = rkme_spec.beta.to(self.device) diff --git a/tests/test_market/test_hetero_market/test_hetero.py b/tests/test_market/test_hetero_market/test_hetero.py index 2d367ac..83d4d95 100644 --- a/tests/test_market/test_hetero_market/test_hetero.py +++ b/tests/test_market/test_hetero_market/test_hetero.py @@ -155,6 +155,7 @@ class TestMarket(unittest.TestCase): hetero_market = self._init_learnware_market() self.test_prepare_learnware_randomly(learnware_num) self.learnware_num = learnware_num + hetero_market.learnware_organizer.reset(auto_update=True, auto_update_limit=learnware_num) print("Total Item:", len(hetero_market)) assert len(hetero_market) == 0, f"The market should be empty!" @@ -173,8 +174,8 @@ class TestMarket(unittest.TestCase): 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}!" - organizer=hetero_market.learnware_organizer - organizer.train() + # organizer=hetero_market.learnware_organizer + # organizer.train(hetero_market.learnware_organizer.learnware_list.values()) return hetero_market def test_search_semantics(self, learnware_num=5):