| @@ -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 | |||
| 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 | |||