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.

methods.py 5.2 kB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105
  1. import numpy as np
  2. from sklearn.metrics import mean_squared_error
  3. from sklearn.model_selection import train_test_split # Add missing import
  4. from loguru import logger
  5. from learnware.reuse import AveragingReuser, EnsemblePruningReuser, FeatureAugmentReuser, HeteroMapAlignLearnware
  6. from examples.dataset_table_workflow.config import align_model_params
  7. def loss_func_rmse(y_true, y_pred):
  8. return np.sqrt(mean_squared_error(y_true, y_pred))
  9. def loss_func_mse(y_true, y_pred):
  10. return mean_squared_error(y_true, y_pred)
  11. def user_model_score(x_train, y_train, test_info):
  12. data_loader = test_info["data_loader"]
  13. x_train, x_val, y_train, y_val = train_test_split(x_train, y_train, test_size=0.2, random_state=42)
  14. user_model = data_loader.train_a_model(x_train, y_train, x_val, y_val)
  15. return user_model
  16. class HomoScoringMethods:
  17. @staticmethod
  18. def single_aug_score(x_train, y_train, test_info):
  19. single_learnware = test_info["single_learnware"]
  20. reuse_single_augment = FeatureAugmentReuser(single_learnware, mode="regression")
  21. reuse_single_augment.fit(x_train=x_train, y_train=y_train)
  22. return reuse_single_augment
  23. @staticmethod
  24. def multiple_aug_score(x_train, y_train, test_info):
  25. multiple_learnwares = test_info["multiple_learnwares"]
  26. reuse_multiple_augment = FeatureAugmentReuser(multiple_learnwares, mode="regression")
  27. reuse_multiple_augment.fit(x_train=x_train, y_train=y_train)
  28. return reuse_multiple_augment
  29. @staticmethod
  30. def multiple_avg_score(x_train, y_train, test_info):
  31. multiple_learnwares = test_info["multiple_learnwares"]
  32. reuse_multiple_avg = AveragingReuser(multiple_learnwares, mode="mean")
  33. return reuse_multiple_avg
  34. @staticmethod
  35. def multiple_ensemble_pruning_score(x_train, y_train, test_info):
  36. multiple_learnwares = test_info["multiple_learnwares"]
  37. if len(multiple_learnwares) == 1:
  38. return multiple_learnwares[0]
  39. reuse_pruning = EnsemblePruningReuser(multiple_learnwares, mode="regression")
  40. reuse_pruning.fit(val_X=x_train, val_y=y_train)
  41. return reuse_pruning
  42. class HeteroMethods:
  43. @staticmethod
  44. def create_hetero_learnware_list(learnware_list, user_rkme, x_train, y_train): # Fix typo in method name
  45. hetero_learnware_list = []
  46. for learnware in learnware_list:
  47. hetero_learnware = HeteroMapAlignLearnware(learnware, mode="regression", **align_model_params)
  48. hetero_learnware.align(user_rkme, x_train, y_train)
  49. hetero_learnware_list.append(hetero_learnware)
  50. return hetero_learnware_list
  51. @staticmethod
  52. def single_aug_score(x_train, y_train, test_info):
  53. user_rkme, single_learnware = test_info["user_rkme"], test_info["single_learnware"]
  54. reuse_single_augment = HeteroMapAlignLearnware(single_learnware, mode="regression", **align_model_params)
  55. reuse_single_augment.align(user_rkme=user_rkme, x_train=x_train, y_train=y_train)
  56. return reuse_single_augment
  57. @staticmethod
  58. def multiple_aug_score(x_train, y_train, test_info):
  59. user_rkme, multiple_learnwares = test_info["user_rkme"], test_info["multiple_learnwares"]
  60. hetero_learnware_list = HeteroMethods.create_hetero_learnware_list(multiple_learnwares, user_rkme, x_train, y_train)
  61. reuse_multiple_augment = FeatureAugmentReuser(hetero_learnware_list, mode="regression")
  62. reuse_multiple_augment.fit(x_train=x_train, y_train=y_train)
  63. return reuse_multiple_augment
  64. @staticmethod
  65. def multiple_ensemble_pruning_score(x_train, y_train, test_info):
  66. user_rkme, multiple_learnwares = test_info["user_rkme"], test_info["multiple_learnwares"]
  67. hetero_learnware_list = HeteroMethods.create_hetero_learnware_list(multiple_learnwares, user_rkme, x_train, y_train)
  68. if len(hetero_learnware_list) == 1:
  69. return hetero_learnware_list[0]
  70. reuse_pruning = EnsemblePruningReuser(hetero_learnware_list, mode="regression")
  71. reuse_pruning.fit(val_X=x_train, val_y=y_train)
  72. return reuse_pruning
  73. @staticmethod
  74. def multiple_avg_score(x_train, y_train, test_info):
  75. user_rkme, multiple_learnwares = test_info["user_rkme"], test_info["multiple_learnwares"]
  76. hetero_learnware_list = HeteroMethods.create_hetero_learnware_list(multiple_learnwares, user_rkme, x_train, y_train)
  77. reuse_multiple_avg = AveragingReuser(hetero_learnware_list, mode="mean")
  78. return reuse_multiple_avg
  79. test_methods = {
  80. "user_model": user_model_score,
  81. "hetero_single_aug": HeteroMethods.single_aug_score,
  82. "hetero_multiple_aug": HeteroMethods.multiple_aug_score,
  83. "hetero_multiple_avg": HeteroMethods.multiple_avg_score,
  84. "hetero_ensemble_pruning": HeteroMethods.multiple_ensemble_pruning_score,
  85. "homo_single_aug": HomoScoringMethods.single_aug_score,
  86. "homo_multiple_aug": HomoScoringMethods.multiple_aug_score,
  87. "homo_multiple_avg": HomoScoringMethods.multiple_avg_score,
  88. "homo_ensemble_pruning": HomoScoringMethods.multiple_ensemble_pruning_score
  89. }