| @@ -43,5 +43,4 @@ cache/ | |||
| tmp/ | |||
| learnware_pool/ | |||
| PFS/ | |||
| data/ | |||
| learnware/market/hetergeneous/.learnware/* | |||
| data/ | |||
| @@ -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 | |||
| @@ -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 | |||
| @@ -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 | |||
| return merged_dfs | |||
| @@ -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) | |||
| @@ -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": { | |||
| @@ -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) | |||
| @@ -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): | |||