You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

base.py 9.5 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199
  1. import os
  2. import joblib
  3. import zipfile
  4. from shutil import copyfile, rmtree
  5. import json
  6. from learnware.client import LearnwareClient
  7. from learnware.logger import get_module_logger
  8. from learnware.market import instantiate_learnware_market
  9. from multiprocessing import Pool
  10. from benchmarks import DataLoader
  11. from config import *
  12. from methods import *
  13. from utils import process_single_aug
  14. logger = get_module_logger("TableWorkflow", level="INFO")
  15. class TableWorkflow:
  16. def __init__(self, learnware_market):
  17. self.learnware_market = learnware_market
  18. self.root_path = os.path.abspath(os.path.join(__file__, ".."))
  19. self.learnware_pool_path = os.path.join(self.root_path, "data/learnware_pool")
  20. self.learnware_zip_pool_path = os.path.join(self.root_path, "data/zips")
  21. self.example_learnware_path = os.path.join(self.root_path, "data/example_files")
  22. self.model_save_path = os.path.join(self.root_path, "data/uploader_models")
  23. self.result_path = os.path.join(self.root_path, "results")
  24. os.makedirs(self.learnware_pool_path, exist_ok=True)
  25. os.makedirs(self.learnware_zip_pool_path, exist_ok=True)
  26. os.makedirs(self.model_save_path, exist_ok=True)
  27. os.makedirs(self.result_path, exist_ok=True)
  28. def _init_dataset(self):
  29. self._prepare_data()
  30. self._prepare_model()
  31. @staticmethod
  32. def _limited_data(method, test_info, loss_func):
  33. all_scores = []
  34. for subset in test_info["train_subsets"]:
  35. subset_scores = []
  36. for sample in subset:
  37. x_train, y_train = sample["x_train"], sample["y_train"]
  38. model = method(x_train, y_train, test_info)
  39. subset_scores.append(loss_func(model.predict(test_info["test_x"]), test_info["test_y"]))
  40. all_scores.append(np.mean(subset_scores))
  41. return all_scores
  42. # @staticmethod
  43. # def _limited_data_single_learnware(method, test_info, learnware):
  44. # test_info['single_learnware'] = learnware
  45. # return TableWorkflow._limited_data(method, test_info)
  46. def test_method(self, test_info, recorders, loss_func=loss_func_rmse):
  47. method_name_full = test_info["method_name"]
  48. method_name = method_name_full if method_name_full == "user_model" else "_".join(method_name_full.split("_")[1:])
  49. user, idx = test_info["user"], test_info["idx"]
  50. recorder = recorders[method_name_full]
  51. save_root_path = os.path.join(self.curves_result_path, f"{user}/{user}_{idx}")
  52. os.makedirs(save_root_path, exist_ok=True)
  53. save_path = os.path.join(save_root_path, f"{method_name}.json")
  54. if method_name == "single_aug":
  55. if test_info["force"] or recorder.should_test_method(user, idx, save_path):
  56. # with Pool() as pool:
  57. # learnware_results = pool.starmap(
  58. # self._limited_data_single_learnware,
  59. # [(test_methods[method_name], test_info, learnware) for learnware in test_info['learnwares']]
  60. # )
  61. # for scores in learnware_results:
  62. # recorders[method_name].record(user, idx, scores)
  63. for learnware in test_info['learnwares']:
  64. test_info['single_learnware'] = learnware
  65. scores = self._limited_data(test_methods[method_name_full], test_info, loss_func)
  66. recorder.record(user, idx, scores)
  67. process_single_aug(user, idx, scores, recorders, save_root_path)
  68. recorder.save(save_path)
  69. logger.info(f"Method {method_name} on {user}_{idx} finished")
  70. else:
  71. process_single_aug(user, idx, recorder.data[user][str(idx)], recorders, save_root_path)
  72. logger.info(f"Method {method_name} on {user}_{idx} already exists")
  73. else:
  74. if test_info["force"] or recorder.should_test_method(user, idx, save_path):
  75. scores = self._limited_data(test_methods[method_name_full], test_info, loss_func)
  76. recorder.record(user, idx, scores)
  77. recorder.save(save_path)
  78. logger.info(f"Method {method_name} on {user}_{idx} finished")
  79. else:
  80. logger.info(f"Method {method_name} on {user}_{idx} already exists")
  81. def prepare_market(self, name, market_id, regenerate_flag=False):
  82. if regenerate_flag:
  83. self._init_dataset()
  84. market = instantiate_learnware_market(name=name, market_id=market_id, rebuild=True)
  85. client = LearnwareClient()
  86. full_descriptions_dir = os.path.join("./data/full_descriptions.json")
  87. with open(full_descriptions_dir, "rb") as f:
  88. full_descriptions = json.load(f)
  89. for uploader in self.learnware_market:
  90. data_loader = DataLoader(uploader)
  91. idx_list = data_loader.get_shop_ids()
  92. for i, idx in enumerate(idx_list):
  93. feature_descriptions = data_loader.get_raw_data(idx)[-1]
  94. feature_dim = len(feature_descriptions)
  95. feature_descriptions_dict = {str(i): feature_descriptions[i] for i in range(feature_dim)}
  96. input_description = {"Dimension": feature_dim, "Description": feature_descriptions_dict}
  97. name_and_description = full_descriptions[uploader][i]
  98. semantic_spec = client.create_semantic_specification(
  99. name=name_and_description["name"],
  100. description=name_and_description["description"],
  101. data_type="Table",
  102. task_type="Regression",
  103. library_type="Others",
  104. license=["MIT"],
  105. scenarios=["Business"],
  106. input_description=input_description,
  107. output_description=output_description,
  108. )
  109. learnware_zip_path = self._prepare_learnware(data_loader, idx)
  110. market.add_learnware(learnware_zip_path, semantic_spec)
  111. # if use pretrained market mapping
  112. if name == "hetero":
  113. learnware_ids = market.get_learnware_ids()
  114. market.learnware_organizer._update_learware_hetero_spec(learnware_ids)
  115. logger.info("Total Item: %d" % (len(market)))
  116. def _prepare_data(self):
  117. for uploader in self.learnware_market:
  118. data_loader = DataLoader(uploader)
  119. data_loader.regenerate_raw_data()
  120. def _prepare_model(self, use_exist=True):
  121. self.learnware_num = 0
  122. for uploader in self.learnware_market:
  123. data_loader = DataLoader(uploader)
  124. idx_list = data_loader.get_shop_ids()
  125. self.learnware_num += len(idx_list)
  126. for idx in idx_list:
  127. logger.info(f"Train on uploader: {uploader}_{idx}")
  128. idx_model_save_path = os.path.join(self.model_save_path, f"{uploader}_{idx}.out")
  129. if not use_exist:
  130. x_train, y_train, x_val, y_val, _ = data_loader.get_raw_data(idx)
  131. data_loader.train_a_model(x_train, y_train, x_val, y_val, save_dir=idx_model_save_path)
  132. else:
  133. uploader_dataset = uploader.split("_")[0]
  134. model = data_loader.get_model(idx)
  135. if uploader_dataset == "corporacion":
  136. model.save_model(idx_model_save_path)
  137. elif uploader_dataset == "pfs":
  138. joblib.dump(model, idx_model_save_path)
  139. else:
  140. logger.error(f"Not supported dataset type {uploader_dataset}")
  141. logger.info(f"Model saved to {idx_model_save_path}")
  142. def _prepare_learnware(self, data_loader, idx):
  143. zip_path = os.path.join(self.learnware_zip_pool_path, f"{data_loader.dataset}_{idx}")
  144. dir_path = os.path.join(self.learnware_pool_path, f"{data_loader.dataset}_{idx}")
  145. model_path = os.path.join(self.model_save_path, f"{data_loader.dataset}_{idx}.out")
  146. os.makedirs(dir_path, exist_ok=True)
  147. stat_spec, _ = data_loader.get_rkme(idx)
  148. init_file = os.path.join(dir_path, "__init__.py")
  149. yaml_file = os.path.join(dir_path, "learnware.yaml")
  150. env_file = os.path.join(dir_path, "environment.yaml")
  151. model_file = os.path.join(dir_path, "model.out")
  152. stat_spec.save(os.path.join(dir_path, "rkme.json"))
  153. copyfile(os.path.join(self.example_learnware_path, f"{data_loader.dataset}/__init__.py"), init_file)
  154. copyfile(os.path.join(self.example_learnware_path, f"{data_loader.dataset}/learnware.yaml"), yaml_file)
  155. copyfile(os.path.join(self.example_learnware_path, "environment.yaml"), env_file)
  156. copyfile(model_path, model_file)
  157. zip_file = zip_path + ".zip"
  158. with zipfile.ZipFile(zip_file, "w") as zip_obj:
  159. for foldername, _, filenames in os.walk(dir_path):
  160. for filename in filenames:
  161. file_path = os.path.join(foldername, filename)
  162. zip_info = zipfile.ZipInfo(filename)
  163. zip_info.compress_type = zipfile.ZIP_STORED
  164. with open(file_path, "rb") as file:
  165. zip_obj.writestr(zip_info, file.read())
  166. rmtree(dir_path) # rm -r dir_path
  167. return zip_file