Browse Source

[ENH] modify add_learnware

tags/v0.3.2
liuht 2 years ago
parent
commit
e65f43c168
1 changed files with 68 additions and 47 deletions
  1. +68
    -47
      learnware/market/heterogeneous/organizer/__init__.py

+ 68
- 47
learnware/market/heterogeneous/organizer/__init__.py View File

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

Loading…
Cancel
Save