diff --git a/learnware/market/heterogeneous/organizer/__init__.py b/learnware/market/heterogeneous/organizer/__init__.py index 9933c19..fdc73fe 100644 --- a/learnware/market/heterogeneous/organizer/__init__.py +++ b/learnware/market/heterogeneous/organizer/__init__.py @@ -37,7 +37,8 @@ class HeteroMapTableOrganizer(EasyOrganizer): self.learnware_zip_list = {} self.learnware_folder_list = {} self.count = 0 - self.last_trained_learnware_num = 0 + self.training_count = 1 + self.last_training_count = 0 self.dbops = DatabaseOperations(conf.database_url, "market_" + self.market_id) self.auto_update = False self.auto_update_limit = auto_update_limit @@ -65,12 +66,25 @@ class HeteroMapTableOrganizer(EasyOrganizer): ) = self.dbops.load_market() if os.path.exists(self.market_mapping_path): - logger.info(f"Loading Market Mapping from Default Checkpoint {self.market_mapping_path}") + logger.info(f"Reload market mapping from checkpoint {self.market_mapping_path}") self.market_mapping = HeteroMapping.load(checkpoint=self.market_store_path) - # self._update_learnware_list(self.learnware_list) + if not rebuild: + if os.path.exists(self.hetero_mappings_path): + for hetero_json_path in os.listdir(self.hetero_mappings_path): + idx = hetero_json_path.split('.')[0] + hetero_spec = HeteroSpecification() + hetero_spec.load(os.path.join(self.hetero_mappings_path, f"{idx}.json")) + try: + self.learnware_list[idx].update_stat_spec("HeteroSpecification", hetero_spec) + except: + logger.warning(f"Learnware ID {idx} NOT Found!") + else: + logger.info("No HeteroSpecifications to reload. Use loaded market mapping to regenerate.") + self._update_learnware_by_ids(self.learnware_list.keys()) else: - logger.warning(f"No Existing Market Mapping!!") + logger.warning(f"No market mapping to reload!!") self.market_mapping = HeteroMapping() + # rmtree(self.hetero_mappings_path) def reset(self, market_id=None, auto_update=False, auto_update_limit=None, **kwargs): self.auto_update = auto_update @@ -81,6 +95,11 @@ class HeteroMapTableOrganizer(EasyOrganizer): def add_learnware( self, zip_path: str, semantic_spec: dict, check_status: int, learnware_id: str = None ) -> Tuple[str, int]: + 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)) learnware_id = "%08d" % (self.count) if learnware_id is None else learnware_id @@ -118,32 +137,34 @@ class HeteroMapTableOrganizer(EasyOrganizer): use_flag=learnwere_status, ) - self._update_learnware_list([new_learnware]) self.learnware_list[learnware_id] = new_learnware 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._update_learnware_by_ids([learnware_id]) self.count += 1 + self.training_count += ([learnware_id] == self._get_table_type_learnware_ids([learnware_id])) - 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()}") + if self.auto_update and self.training_count - self.last_training_count == self.auto_update_limit + 1: + training_learnware_ids = self._get_table_type_learnware_ids(self.get_learnware_ids()) + training_learnwares = self.get_learnware_by_ids(training_learnware_ids) + logger.warning(f"Leanwares for training: {training_learnware_ids}") updated_market_mapping = self.train( - learnware_list=self.learnware_list.values(), + learnware_list=training_learnwares, save_dir=self.market_store_path, **self.training_args ) - logger.warning(f"Market mapping train completed. Now update HeteroSpecification for {self.get_learnware_ids()}") - + logger.warning(f"Market mapping train completed. Now update HeteroSpecification for {training_learnware_ids}") self.market_mapping = updated_market_mapping - self._update_learnware_list(self.learnware_list.values()) - self.last_trained_learnware_num = self.count + self._update_learnware_by_ids(training_learnware_ids) + self.last_training_count = len(training_learnware_ids) return learnware_id, learnwere_status @staticmethod - def train(learnware_list: List[Learnware] = None, save_dir: str = None, **kwargs): + def train(learnware_list: List[Learnware], save_dir: str, **kwargs): allset = HeteroMapTableOrganizer._learnwares_to_dataframes(learnware_list) market_mapping = HeteroMapping(**kwargs) market_mapping_trainer = Trainer( @@ -157,45 +178,45 @@ class HeteroMapTableOrganizer(EasyOrganizer): market_mapping_trainer.save_model(output_dir=save_dir) return market_mapping - - ############################################ - # save_model & generateing new specification - # should be moved out of train thread - ############################################ - - def _update_learnware_list(self, learnware_list: List[Learnware]): - try: - for learnware in learnware_list: - 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 HeteroSpecification failed! Due to {err}") - - def _update_learnware_specification(self, learnware: Learnware, save_path: str) -> Learnware: - specification = learnware.specification - learnware_rkme = specification.get_stat_spec()["RKMETableSpecification"] - learnware_features = specification.get_semantic_spec()["Input"]["Description"].values() - learnware_hetero_spec = self.market_mapping.hetero_mapping(learnware_rkme, learnware_features) - learnware.update_stat_spec("HeteroSpecification", learnware_hetero_spec) - - learnware_hetero_spec.save(save_path) - + + def _update_learnware_by_ids(self, ids: List[str]): + ids = self._get_table_type_learnware_ids(ids) + for id in ids: + try: + spec = self.learnware_list[id].get_specification() + semantic_spec, stat_spec = spec.get_semantic_spec(), spec.get_stat_spec()["RKMETableSpecification"] + features = semantic_spec["Input"]["Description"].values() + hetero_spec = self.market_mapping.hetero_mapping(stat_spec, features) + self.learnware_list[id].update_stat_spec("HeteroSpecification", hetero_spec) + + save_path = os.path.join(self.hetero_mappings_path, f"{id}.json") + hetero_spec.save(save_path) + except Exception as err: + logger.warning(f"Learnware {id} generate HeteroSpecification failed! Due to {err}") + def generate_hetero_map_spec(self, user_info: BaseUserInfo) -> HeteroSpecification: - user_rkme = user_info.stat_info["RKMETableSpecification"] + user_stat_spec = user_info.stat_info["RKMETableSpecification"] user_features = user_info.get_semantic_spec()["Input"]["Description"].values() - user_hetero_spec = self.market_mapping.hetero_mapping(user_rkme, user_features) + + user_hetero_spec = self.market_mapping.hetero_mapping(user_stat_spec, user_features) return user_hetero_spec @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() - learnware_rkme = specification.get_stat_spec()["RKMETableSpecification"] - learnware_features = specification.get_semantic_spec()["Input"]["Description"] - learnware_df = pd.DataFrame(data=learnware_rkme.get_z(), columns=learnware_features.values()) - - 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 \ No newline at end of file + spec = learnware.get_specification() + stat_spec = spec.get_stat_spec()["RKMETableSpecification"] + features = spec.get_semantic_spec()["Input"]["Description"] + learnware_df = pd.DataFrame(data=stat_spec.get_z(), columns=features.values()) + learnware_df_dict[tuple(sorted(features))].append(learnware_df) + + return [pd.concat(dfs) for dfs in learnware_df_dict.values()] + + def _get_table_type_learnware_ids(self, ids: List[str]) -> List[str]: + ret = [] + for id in ids: + semantic_spec = self.learnware_list[id].get_specification().get_semantic_spec() + if semantic_spec["Data"]["Values"][0] == "Table": + ret.append(id) + return ret \ No newline at end of file