PyIAML 1.0.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (279) hide show
  1. iaml/__init__.py +56 -0
  2. iaml/actionable.py +11 -0
  3. iaml/actionables/__init__.py +21 -0
  4. iaml/actionables/boosting/__init__.py +4 -0
  5. iaml/actionables/boosting/act_adaboost.py +59 -0
  6. iaml/actionables/cleaning/__init__.py +26 -0
  7. iaml/actionables/cleaning/act_categorical_imputer.py +124 -0
  8. iaml/actionables/cleaning/act_count_vectorizer.py +204 -0
  9. iaml/actionables/cleaning/act_drop_categorical_column.py +51 -0
  10. iaml/actionables/cleaning/act_drop_date_column.py +48 -0
  11. iaml/actionables/cleaning/act_drop_high_cardinality_categorical.py +337 -0
  12. iaml/actionables/cleaning/act_drop_numerical_column.py +75 -0
  13. iaml/actionables/cleaning/act_drop_textual_column.py +51 -0
  14. iaml/actionables/cleaning/act_encode_target_column.py +56 -0
  15. iaml/actionables/cleaning/act_frequency_encoder.py +127 -0
  16. iaml/actionables/cleaning/act_hashing_vectorizer.py +186 -0
  17. iaml/actionables/cleaning/act_knn_imputer.py +152 -0
  18. iaml/actionables/cleaning/act_mean_column.py +79 -0
  19. iaml/actionables/cleaning/act_mice.py +464 -0
  20. iaml/actionables/cleaning/act_missing_count_feature.py +109 -0
  21. iaml/actionables/cleaning/act_missing_indicator.py +124 -0
  22. iaml/actionables/cleaning/act_onehot.py +65 -0
  23. iaml/actionables/cleaning/act_ordinal_encoder.py +177 -0
  24. iaml/actionables/cleaning/act_rare_category_grouper.py +173 -0
  25. iaml/actionables/cleaning/act_simple_imputer.py +109 -0
  26. iaml/actionables/cleaning/act_split_date.py +68 -0
  27. iaml/actionables/cleaning/act_target_encoder.py +274 -0
  28. iaml/actionables/cleaning/act_text_normalizer.py +241 -0
  29. iaml/actionables/cleaning/act_tf_idf.py +80 -0
  30. iaml/actionables/cleaning/act_word2vec.py +150 -0
  31. iaml/actionables/features_precleaning/__init__.py +12 -0
  32. iaml/actionables/features_precleaning/act_coerce_numeric_strings.py +194 -0
  33. iaml/actionables/features_precleaning/act_date_converter.py +99 -0
  34. iaml/actionables/features_precleaning/act_drop_bad_quality_rows.py +77 -0
  35. iaml/actionables/features_precleaning/act_drop_duplicate_rows.py +131 -0
  36. iaml/actionables/features_precleaning/act_drop_high_missing_columns.py +94 -0
  37. iaml/actionables/features_precleaning/act_drop_id_like_columns.py +294 -0
  38. iaml/actionables/features_precleaning/act_normalize_column_names.py +157 -0
  39. iaml/actionables/features_precleaning/act_sentinel_to_na_n.py +270 -0
  40. iaml/actionables/features_precleaning/act_trim_space.py +79 -0
  41. iaml/actionables/features_preprocessing/__init__.py +18 -0
  42. iaml/actionables/features_preprocessing/act_cyclical_date_encoding.py +212 -0
  43. iaml/actionables/features_preprocessing/act_fast_ica.py +161 -0
  44. iaml/actionables/features_preprocessing/act_feature_agglomeration.py +90 -0
  45. iaml/actionables/features_preprocessing/act_k_bins_discretizer.py +207 -0
  46. iaml/actionables/features_preprocessing/act_k_means_features.py +296 -0
  47. iaml/actionables/features_preprocessing/act_kernel_pca.py +143 -0
  48. iaml/actionables/features_preprocessing/act_log_transformer.py +122 -0
  49. iaml/actionables/features_preprocessing/act_nystroem.py +100 -0
  50. iaml/actionables/features_preprocessing/act_pca.py +77 -0
  51. iaml/actionables/features_preprocessing/act_polynomial_features.py +86 -0
  52. iaml/actionables/features_preprocessing/act_power_transformer.py +106 -0
  53. iaml/actionables/features_preprocessing/act_quantile_transformer.py +114 -0
  54. iaml/actionables/features_preprocessing/act_rbf_sampler.py +88 -0
  55. iaml/actionables/features_preprocessing/act_select_percentile.py +112 -0
  56. iaml/actionables/features_preprocessing/act_sparse_random_projection.py +157 -0
  57. iaml/actionables/features_preprocessing/act_truncated_svd.py +137 -0
  58. iaml/actionables/features_selection/__init__.py +8 -0
  59. iaml/actionables/features_selection/act_permutation_importance_selector.py +421 -0
  60. iaml/actionables/features_selection/act_remove_high_correlated_column.py +70 -0
  61. iaml/actionables/features_selection/act_remove_low_variance_column.py +74 -0
  62. iaml/actionables/features_selection/act_rfe.py +214 -0
  63. iaml/actionables/features_selection/act_select_from_model.py +325 -0
  64. iaml/actionables/features_selection/act_select_k_best.py +181 -0
  65. iaml/actionables/features_selection/act_vif_selector.py +130 -0
  66. iaml/actionables/imbalance/__init__.py +10 -0
  67. iaml/actionables/imbalance/act_adasyn.py +150 -0
  68. iaml/actionables/imbalance/act_borderline_smote.py +171 -0
  69. iaml/actionables/imbalance/act_near_miss.py +158 -0
  70. iaml/actionables/imbalance/act_random_over_sampling.py +60 -0
  71. iaml/actionables/imbalance/act_random_under_sampler.py +135 -0
  72. iaml/actionables/imbalance/act_smote.py +162 -0
  73. iaml/actionables/imbalance/act_smote_tomek.py +182 -0
  74. iaml/actionables/imbalance/act_smoteenn.py +193 -0
  75. iaml/actionables/imbalance/act_tomek_links.py +138 -0
  76. iaml/actionables/normalize/__init__.py +6 -0
  77. iaml/actionables/normalize/act_max_abs_scaler.py +78 -0
  78. iaml/actionables/normalize/act_minmax_scaler.py +56 -0
  79. iaml/actionables/normalize/act_normalizer.py +95 -0
  80. iaml/actionables/normalize/act_robust_scaler.py +111 -0
  81. iaml/actionables/normalize/act_standard_scaler.py +55 -0
  82. iaml/actionables/predictors/__init__.py +6 -0
  83. iaml/actionables/predictors/_xgboost.py +16 -0
  84. iaml/actionables/predictors/classifier/__init__.py +26 -0
  85. iaml/actionables/predictors/classifier/act_bagging_classifier.py +113 -0
  86. iaml/actionables/predictors/classifier/act_bernoulli_nb.py +89 -0
  87. iaml/actionables/predictors/classifier/act_catboost_classifier.py +135 -0
  88. iaml/actionables/predictors/classifier/act_complement_nb.py +106 -0
  89. iaml/actionables/predictors/classifier/act_decision_tree_classifier.py +117 -0
  90. iaml/actionables/predictors/classifier/act_extra_trees_classifier.py +115 -0
  91. iaml/actionables/predictors/classifier/act_gaussian_nb.py +53 -0
  92. iaml/actionables/predictors/classifier/act_hist_gradient_boosting_classifier.py +144 -0
  93. iaml/actionables/predictors/classifier/act_knn.py +86 -0
  94. iaml/actionables/predictors/classifier/act_light_gbm_classifier.py +211 -0
  95. iaml/actionables/predictors/classifier/act_linear_discriminant_analysis.py +63 -0
  96. iaml/actionables/predictors/classifier/act_linear_svc.py +134 -0
  97. iaml/actionables/predictors/classifier/act_logistic_regression.py +92 -0
  98. iaml/actionables/predictors/classifier/act_mlp_classifier.py +107 -0
  99. iaml/actionables/predictors/classifier/act_multinomial_nb.py +76 -0
  100. iaml/actionables/predictors/classifier/act_passive_aggressive_classifier.py +141 -0
  101. iaml/actionables/predictors/classifier/act_quadratic_discriminant_analysis.py +72 -0
  102. iaml/actionables/predictors/classifier/act_randomforest.py +113 -0
  103. iaml/actionables/predictors/classifier/act_ridge_classifier.py +116 -0
  104. iaml/actionables/predictors/classifier/act_sgd_classifier.py +149 -0
  105. iaml/actionables/predictors/classifier/act_svm_svc.py +88 -0
  106. iaml/actionables/predictors/classifier/act_xgboost.py +111 -0
  107. iaml/actionables/predictors/regressor/__init__.py +27 -0
  108. iaml/actionables/predictors/regressor/act_ada_boost_regressor.py +75 -0
  109. iaml/actionables/predictors/regressor/act_ard_regression.py +95 -0
  110. iaml/actionables/predictors/regressor/act_catboost_regressor.py +134 -0
  111. iaml/actionables/predictors/regressor/act_decision_tree_regressor.py +111 -0
  112. iaml/actionables/predictors/regressor/act_elastic_net_regressor.py +109 -0
  113. iaml/actionables/predictors/regressor/act_extra_trees_regressor.py +113 -0
  114. iaml/actionables/predictors/regressor/act_gaussian_process_regressor.py +55 -0
  115. iaml/actionables/predictors/regressor/act_gboost_regressor.py +95 -0
  116. iaml/actionables/predictors/regressor/act_hist_gradient_boosting_regressor.py +105 -0
  117. iaml/actionables/predictors/regressor/act_huber_regressor.py +101 -0
  118. iaml/actionables/predictors/regressor/act_knn_regressor.py +86 -0
  119. iaml/actionables/predictors/regressor/act_lasso_regressor.py +103 -0
  120. iaml/actionables/predictors/regressor/act_light_gbm_regressor.py +201 -0
  121. iaml/actionables/predictors/regressor/act_linear_regression.py +43 -0
  122. iaml/actionables/predictors/regressor/act_mlp_regressor.py +104 -0
  123. iaml/actionables/predictors/regressor/act_poisson_regressor.py +111 -0
  124. iaml/actionables/predictors/regressor/act_quantile_regressor.py +87 -0
  125. iaml/actionables/predictors/regressor/act_randomforest_regressor.py +116 -0
  126. iaml/actionables/predictors/regressor/act_ransac_regressor.py +106 -0
  127. iaml/actionables/predictors/regressor/act_ridge_regressor.py +107 -0
  128. iaml/actionables/predictors/regressor/act_sgd_regressor.py +106 -0
  129. iaml/actionables/predictors/regressor/act_svm_svr.py +81 -0
  130. iaml/actionables/predictors/regressor/act_xgboost_regressor.py +97 -0
  131. iaml/actionables/predictors/survival/__init__.py +12 -0
  132. iaml/actionables/predictors/survival/act_aalen_additive_model.py +83 -0
  133. iaml/actionables/predictors/survival/act_cox.py +110 -0
  134. iaml/actionables/predictors/survival/act_coxnet_survival_analysis.py +134 -0
  135. iaml/actionables/predictors/survival/act_extra_survival_trees.py +101 -0
  136. iaml/actionables/predictors/survival/act_fast_survival_svm.py +102 -0
  137. iaml/actionables/predictors/survival/act_gradient_boosting_survival_analysis.py +93 -0
  138. iaml/actionables/predictors/survival/act_random_survival_forest.py +91 -0
  139. iaml/actionables/predictors/survival/act_survival_component_wise_gboost.py +80 -0
  140. iaml/actionables/predictors/survival/act_survival_tree.py +120 -0
  141. iaml/actionables/predictors/survival/act_survival_xgboost.py +9 -0
  142. iaml/actionables/predictors/survival/act_weibull_aft.py +230 -0
  143. iaml/cache.py +61 -0
  144. iaml/cache_keys.py +57 -0
  145. iaml/candidate.py +736 -0
  146. iaml/core_dispatcher.py +125 -0
  147. iaml/data_type.py +11 -0
  148. iaml/dataset.py +506 -0
  149. iaml/decorators/__init__.py +3 -0
  150. iaml/decorators/all.py +4 -0
  151. iaml/decorators/is_step.py +45 -0
  152. iaml/decorators/runner.py +100 -0
  153. iaml/explanation.py +112 -0
  154. iaml/iaml.py +1072 -0
  155. iaml/iaml_pipeline.py +600 -0
  156. iaml/logger.py +138 -0
  157. iaml/meta_explorer_step.py +62 -0
  158. iaml/meta_ordered_step.py +28 -0
  159. iaml/meta_partial_explorer_step.py +34 -0
  160. iaml/meta_singleton.py +24 -0
  161. iaml/metastep.py +211 -0
  162. iaml/metric.py +111 -0
  163. iaml/metric_plot.py +82 -0
  164. iaml/metrics/__init__.py +21 -0
  165. iaml/metrics/_classification.py +28 -0
  166. iaml/metrics/_survival_times.py +22 -0
  167. iaml/metrics/accuracy_metric.py +59 -0
  168. iaml/metrics/balanced_accuracy_metric.py +67 -0
  169. iaml/metrics/brier_score.py +90 -0
  170. iaml/metrics/classification_error_metric.py +66 -0
  171. iaml/metrics/concordance_index_ipcw.py +84 -0
  172. iaml/metrics/concordance_index_metric.py +67 -0
  173. iaml/metrics/cumulative_dynamic_auc.py +119 -0
  174. iaml/metrics/f1_score_metric.py +71 -0
  175. iaml/metrics/integrated_brier_score.py +98 -0
  176. iaml/metrics/integrated_brier_score_loss.py +41 -0
  177. iaml/metrics/mean_absolute_error_metric.py +46 -0
  178. iaml/metrics/mean_squared_error_metric.py +46 -0
  179. iaml/metrics/mean_squared_log_error_metric.py +49 -0
  180. iaml/metrics/median_absolute_error_metric.py +48 -0
  181. iaml/metrics/precision_metric.py +63 -0
  182. iaml/metrics/r2_score_metric.py +45 -0
  183. iaml/metrics/recall_metric.py +65 -0
  184. iaml/metrics/roc_auc_metric.py +50 -0
  185. iaml/metrics/specificity_metric.py +44 -0
  186. iaml/metrics/specificity_multiclass_metric.py +55 -0
  187. iaml/metrics/specificity_multilabel_metric.py +60 -0
  188. iaml/optimizers/__init__.py +5 -0
  189. iaml/optimizers/bayesian_optimizer.py +193 -0
  190. iaml/optimizers/genetic_optimizer.py +284 -0
  191. iaml/optimizers/optimizer.py +31 -0
  192. iaml/optimizers/random_optimizer.py +101 -0
  193. iaml/plot.py +138 -0
  194. iaml/plots/__init__.py +32 -0
  195. iaml/plots/bar_plot.py +141 -0
  196. iaml/plots/box_plot.py +166 -0
  197. iaml/plots/class_prediction_error_plot.py +37 -0
  198. iaml/plots/classification_report_plot.py +35 -0
  199. iaml/plots/confusion_matrix_plot.py +34 -0
  200. iaml/plots/correlation_heatmap_plot.py +201 -0
  201. iaml/plots/cumulative_hazard_plot.py +72 -0
  202. iaml/plots/density_plot.py +210 -0
  203. iaml/plots/histogram_plot.py +179 -0
  204. iaml/plots/kaplan_meier_comparison_plot.py +89 -0
  205. iaml/plots/line_plot.py +70 -0
  206. iaml/plots/missingness_heatmap_plot.py +203 -0
  207. iaml/plots/outlier_plot.py +217 -0
  208. iaml/plots/pair_plot.py +228 -0
  209. iaml/plots/precision_recall_curve_plot.py +86 -0
  210. iaml/plots/prediction_error_plot.py +34 -0
  211. iaml/plots/qq_plot.py +220 -0
  212. iaml/plots/residual_plot.py +38 -0
  213. iaml/plots/roc_dynamique_curve_plot.py +79 -0
  214. iaml/plots/rocauc_plot.py +96 -0
  215. iaml/plots/shap_plot.py +187 -0
  216. iaml/plots/target_distribution_plot.py +241 -0
  217. iaml/plots/violin_plot.py +206 -0
  218. iaml/predictor.py +139 -0
  219. iaml/reference.py +65 -0
  220. iaml/shared_cache.py +90 -0
  221. iaml/sklearn_preprocessor.py +74 -0
  222. iaml/splitters/__init__.py +3 -0
  223. iaml/splitters/kfold_splitter.py +32 -0
  224. iaml/splitters/random_splitter.py +26 -0
  225. iaml/stack.py +39 -0
  226. iaml/statistic.py +66 -0
  227. iaml/statistics/__init__.py +77 -0
  228. iaml/statistics/anova_statistic.py +80 -0
  229. iaml/statistics/cardinality_ratio_statistic.py +63 -0
  230. iaml/statistics/category_cooccurrence_statistic.py +79 -0
  231. iaml/statistics/chi_square_statistic.py +81 -0
  232. iaml/statistics/coef_variation_statistic.py +72 -0
  233. iaml/statistics/correlation_with_target.py +105 -0
  234. iaml/statistics/count.py +72 -0
  235. iaml/statistics/data_type_summary_statistic.py +74 -0
  236. iaml/statistics/duplicate_row_statistic.py +56 -0
  237. iaml/statistics/effect_size_statistic.py +129 -0
  238. iaml/statistics/entropy_statistic.py +69 -0
  239. iaml/statistics/event_rate_statistic.py +52 -0
  240. iaml/statistics/grouped_mean_statistic.py +60 -0
  241. iaml/statistics/iqr_statistic.py +66 -0
  242. iaml/statistics/kurtosis.py +50 -0
  243. iaml/statistics/mad_statistic.py +66 -0
  244. iaml/statistics/mean.py +61 -0
  245. iaml/statistics/median_statistic.py +61 -0
  246. iaml/statistics/minmax.py +60 -0
  247. iaml/statistics/missing_rate_statistic.py +62 -0
  248. iaml/statistics/mode.py +47 -0
  249. iaml/statistics/most_frequent_ratio.py +81 -0
  250. iaml/statistics/outlier_count_iqr_statistic.py +76 -0
  251. iaml/statistics/quantile.py +59 -0
  252. iaml/statistics/range.py +53 -0
  253. iaml/statistics/rare_category_rate.py +92 -0
  254. iaml/statistics/skewness.py +53 -0
  255. iaml/statistics/stdev.py +50 -0
  256. iaml/statistics/summary_table_statistic.py +60 -0
  257. iaml/statistics/time_by_group_statistic.py +83 -0
  258. iaml/statistics/time_summary_statistic.py +56 -0
  259. iaml/statistics/top_k_value_counts.py +68 -0
  260. iaml/statistics/unique_count_statistic.py +57 -0
  261. iaml/statistics/value_counts.py +63 -0
  262. iaml/statistics/variance.py +51 -0
  263. iaml/statistics/violin.py +63 -0
  264. iaml/step.py +600 -0
  265. iaml/step_cache.py +87 -0
  266. iaml/step_wrapper.py +79 -0
  267. iaml/timed_pool_executor.py +492 -0
  268. iaml/type_of_target.py +68 -0
  269. iaml/void_step.py +101 -0
  270. iaml/worker_manager.py +169 -0
  271. iaml/wrapper/__init__.py +4 -0
  272. iaml/wrapper/wrap_basic_gridsearch.py +68 -0
  273. iaml/wrapper/wrap_genetic_gridsearch.py +293 -0
  274. iaml/wrapper/wrap_iterative_gridsearch.py +399 -0
  275. pyiaml-1.0.0.dist-info/METADATA +802 -0
  276. pyiaml-1.0.0.dist-info/RECORD +279 -0
  277. pyiaml-1.0.0.dist-info/WHEEL +5 -0
  278. pyiaml-1.0.0.dist-info/licenses/LICENSE +674 -0
  279. pyiaml-1.0.0.dist-info/top_level.txt +1 -0
iaml/iaml_pipeline.py ADDED
@@ -0,0 +1,600 @@
1
+ """Based on Scikit-learn Pipeline but for IAML Pipelines !
2
+ Apply preprocessing in construction order, then predict from Candidate instance.
3
+ """
4
+ from __future__ import annotations
5
+ from typing import TYPE_CHECKING
6
+
7
+ import pickle
8
+
9
+ from copy import deepcopy
10
+ from hashlib import md5
11
+
12
+ import numpy as np
13
+ import shap
14
+ import pandas as pd
15
+
16
+ from sklearn.pipeline import Pipeline
17
+
18
+ from .dataset import Dataset
19
+ from .void_step import VoidStep
20
+ from .explanation import Explanation
21
+ from .cache import Cache
22
+ from .reference import Reference
23
+
24
+ if TYPE_CHECKING:
25
+ from .metric import Metric
26
+ from .step import Step
27
+
28
+ class IAMLPipeline(Pipeline):
29
+ """Based on Scikit-learn Pipeline but for IAML Pipelines !
30
+ Apply preprocessing in construction order, then predict from Candidate instance.
31
+
32
+ :param list[tuple[str, Step]], optional steps: Ordered list of IAML.Steps. Defaults to None.
33
+ :param pd.DataFrame, optional original_dataset: Untransformed dataset to use as a masker
34
+ for the SHAP explainer which will be used to explain the model later on.
35
+ Defaults to None. If not provided, the prediction dataset will be used as the
36
+ masker, which may impact the accuracy of the explanations.
37
+ :param str, optional estimator_type: Type of estimator. Must be one of 'classifier',
38
+ 'survival', 'regressor'.
39
+ """
40
+
41
+ def __init__(
42
+ self,
43
+ steps: list[tuple[str, Step]] = None,
44
+ original_dataset: pd.DataFrame = None,
45
+ estimator_type: str = None) -> None:
46
+ if steps is None:
47
+ steps = []
48
+
49
+ self.original_dataset: pd.DataFrame = original_dataset
50
+ """Original dataset used for this pipeline"""
51
+
52
+ self._preprocessing_steps: list[tuple[str, object]] = []
53
+ """Transformers and resamplers in their construction order."""
54
+
55
+ self.predictor: tuple[str, object] = None
56
+ """Predictor that'll be used in this pipeline"""
57
+
58
+ self.metrics: list[Metric] = []
59
+ """List of metrics that'll be computed in this pipeline"""
60
+
61
+ self._fingerprint_cache: str | None = None
62
+ """Cached fingerprint of training steps."""
63
+
64
+ self._transform_fingerprint_cache: str | None = None
65
+ """Cached fingerprint of transforms/resamplers."""
66
+
67
+ self._fingerprint_cache_version: tuple | None = None
68
+ """Cached config versions for training steps."""
69
+
70
+ self._transform_fingerprint_cache_version: tuple | None = None
71
+ """Cached config versions for transforms/resamplers."""
72
+
73
+ self._trained_columns: list[str] | None = None
74
+ """Columns seen by the predictor during fit (after transforms)."""
75
+
76
+ if estimator_type not in ['classifier', 'regressor', 'survival']:
77
+ raise ValueError(f"Estimator type ({estimator_type}) must be classifier, \
78
+ survival or regressor")
79
+ self.__estimator_type: str = estimator_type
80
+ """Estimator type"""
81
+
82
+ super().__init__(steps)
83
+
84
+ @property
85
+ def _estimator_type(self) -> str:
86
+ """Expose the estimator type to scikit-learn.
87
+
88
+ :return: estimator type.
89
+ """
90
+ return self.__estimator_type
91
+
92
+ @property
93
+ def estimator_type(self) -> str:
94
+ """Return the estimator type configured for this pipeline.
95
+
96
+ :return: estimator type.
97
+ """
98
+ return self.__estimator_type
99
+
100
+ @property
101
+ def transformers(self) -> list[tuple[str, object]]:
102
+ """Transform steps in execution order, excluding training-only resamplers."""
103
+ return [step for step in self._preprocessing_steps
104
+ if callable(getattr(step[1], 'transform', None))]
105
+
106
+ @property
107
+ def resamplers(self) -> list[tuple[str, object]]:
108
+ """Resampling steps in execution order."""
109
+ return [step for step in self._preprocessing_steps
110
+ if not callable(getattr(step[1], 'transform', None))
111
+ and callable(getattr(step[1], 'resample', None))]
112
+
113
+ @property
114
+ def steps(self) -> list[tuple[str, object]]:
115
+ """Get prediction steps, excluding training-only resamplers.
116
+
117
+ :return: list of steps
118
+ """
119
+ return [item for item in [*self.transformers, self.predictor] if item is not None]
120
+
121
+ @property
122
+ def training_steps(self) -> list[tuple[str, object]]:
123
+ """Steps used to fit the pipeline, preserving preprocessing order.
124
+
125
+ :return: list of steps
126
+ """
127
+ return [item for item in [*self._preprocessing_steps, self.predictor] \
128
+ if item is not None]
129
+
130
+ @steps.setter
131
+ def steps(self, values: list[tuple[str, object]]) -> list[tuple[str, object]]:
132
+ """Set pipeline steps
133
+
134
+ :param list[tuple[str, object]] values: List of steps to add.
135
+ :return: List of new steps.
136
+ """
137
+ self._preprocessing_steps = []
138
+ self.predictor = None
139
+ self._invalidate_fingerprint_cache()
140
+
141
+ for value in values:
142
+ self.__add_step(value)
143
+
144
+ return self.steps
145
+
146
+ def __add_step(self, step: tuple[str, object]) -> None:
147
+ """Add a step to the pipeline steps
148
+
149
+ :param tuple[str,object] step: The step to add.
150
+ """
151
+ _, instance = step
152
+ self._invalidate_fingerprint_cache()
153
+ if hasattr(instance, 'predict') and callable(instance.predict):
154
+ self.predictor = step
155
+ elif callable(getattr(instance, 'transform', None)) \
156
+ or callable(getattr(instance, 'resample', None)):
157
+ self._preprocessing_steps.append(step)
158
+
159
+ def replace_step(self, old: 'Step', new: 'Step') -> bool:
160
+ """Replace a step in the pipeline by another (by id)
161
+
162
+ :param Step old: The step to replace.
163
+ :param Step new: The new step.
164
+ :return: Was replaced ?
165
+ """
166
+ for idx, step in enumerate(self._preprocessing_steps):
167
+ if old is step[1]:
168
+ self._preprocessing_steps[idx] = (new.name, new)
169
+ self._invalidate_fingerprint_cache()
170
+ return True
171
+ if self.predictor is not None and old is self.predictor[1]:
172
+ self.predictor = (new.name, new)
173
+ self._invalidate_fingerprint_cache()
174
+ return True
175
+
176
+ return False
177
+
178
+ def remove_step(self, to_remove: 'Step') -> bool:
179
+ """Remove a step from the pipeline (by object id)
180
+
181
+ :param Step to_remove: Step to remove.
182
+ :return: Step was removed ?
183
+ """
184
+ for idx, step in enumerate(self._preprocessing_steps):
185
+ if to_remove is step[1]:
186
+ del self._preprocessing_steps[idx]
187
+ self._invalidate_fingerprint_cache()
188
+ return True
189
+ if self.predictor is not None and to_remove is self.predictor[1]:
190
+ self.predictor = None
191
+ self._invalidate_fingerprint_cache()
192
+ return True
193
+
194
+ return False
195
+
196
+ def fit(
197
+ self,
198
+ X: pd.DataFrame,
199
+ y: pd.DataFrame = None,
200
+ only_predictor: bool = False,
201
+ groups_columns: list[str] = None,
202
+ metrics: list[Metric] = None,
203
+ **kwargs: dict) -> 'IAMLPipeline':
204
+ """Fit Pipeline on new data (or with new parameters)
205
+
206
+ :param pd.DataFrame X: Candidate features.
207
+ :param pd.DataFrame, optional y: label to predict. Default to None.
208
+ :param bool, optional only_predictor: Run predictions only. Default to False.
209
+ :param list[str], optional groups_columns: Columns name to use in splitting.
210
+ Default to None.
211
+ :param list[Metric], optional metrics: List of Metrics to compute. Default to None.
212
+ :param dict, optional \\**kwargs: Additional parameters.
213
+ :return: Fitted IAMLPipeline.
214
+ """
215
+ self.metrics = metrics
216
+
217
+ if groups_columns is None:
218
+ groups_columns = []
219
+
220
+ if not only_predictor:
221
+ X, y = self.fit_transform(X, y, groups_columns=groups_columns, **kwargs)
222
+ # Reset groups_columns as returned X is aldready pruned from groups columns
223
+ # This way we avoid caching KeyError in dataset init
224
+ groups_columns = []
225
+ dataset = Dataset(X, y, groups_columns=groups_columns)
226
+
227
+ if self.predictor[1].suitable(dataset):
228
+ if isinstance(dataset.X, pd.DataFrame):
229
+ self._trained_columns = list(dataset.X.columns)
230
+ self.predictor[1].fit(dataset, **kwargs)
231
+ else:
232
+ self.predictor = None
233
+
234
+ return self
235
+
236
+ def fit_transform(
237
+ self,
238
+ X: pd.DataFrame,
239
+ y: pd.DataFrame = None,
240
+ groups_columns: list[str] = None,
241
+ **kwargs: dict) -> 'IAMLPipeline':
242
+ """Fit Pipeline and transform data
243
+
244
+ :param pd.DataFrame X: Candidate features.
245
+ :param pd.DataFrame, optional y: Label to predict. Default to None.
246
+ :param list[str], optional groups_columns: Columns name to use in splitting.
247
+ Default to None.
248
+ :param dict, optional \\**kwargs: Additional parameters.
249
+ :return: Fitted IAMLPipeline.
250
+ """
251
+ if groups_columns is None:
252
+ groups_columns = []
253
+
254
+ dataset = Dataset(X, y, groups_columns=groups_columns)
255
+
256
+ # Cached fits can replace steps, and unsuitable steps can be removed.
257
+ for _, step in list(self._preprocessing_steps):
258
+ # FIT
259
+ if 'Step' in map(lambda s: s.__name__, step.__class__.__mro__):
260
+ fit_key = f"fit_{step.fingerprint()}"
261
+ fit_data_key = dataset.fingerprint()
262
+ from_cache = Cache().from_cache(fit_key, fit_data_key)
263
+ if from_cache is not None:
264
+ fitted_step, dataset = from_cache
265
+ self.replace_step(step, fitted_step)
266
+ step = fitted_step
267
+ else:
268
+ if step.suitable(dataset):
269
+ step.fit(dataset)
270
+ # Copy together to preserve shared references to training data
271
+ # (e.g. a target encoder's out-of-fold training transform).
272
+ Cache().add_to_cache(fit_key, fit_data_key, (step, dataset))
273
+ else:
274
+ if step.is_interchangeable:
275
+ old_step = step
276
+ step = VoidStep(step_to_mimic=step)
277
+ self.replace_step(old_step, step)
278
+ else:
279
+ self.remove_step(step)
280
+ continue
281
+ else:
282
+ step.fit(dataset.X, dataset.y, **kwargs)
283
+
284
+ # APPLY TRANSFORM / RESAMPLE
285
+ apply_key = f"apply_{step.fingerprint()}"
286
+ # Freeze before a step can mutate X or y in place.
287
+ apply_data_key = dataset.fingerprint()
288
+ dataset_from_cache = Cache().from_cache(apply_key, apply_data_key)
289
+
290
+ if dataset_from_cache is not None:
291
+ dataset = dataset_from_cache
292
+ else:
293
+ if hasattr(step, 'transform'):
294
+ dataset.transform(step.transform)
295
+ elif hasattr(step, 'resample'):
296
+ dataset = dataset.resample(step.resample)
297
+ Cache().add_to_cache(apply_key, apply_data_key, dataset)
298
+
299
+ return dataset.X, dataset.y
300
+
301
+ @property
302
+ def explanations(self) -> list[str]:
303
+ """Get explanations from all pipeline steps
304
+
305
+ :return: List of markdown explanations
306
+ """
307
+ return [ e for _, step in self.training_steps if (e := step.explain()) is not None ]
308
+
309
+ @property
310
+ def model(self) -> Step:
311
+ """Shortcut to get the prediction model of IAMLPipeline
312
+
313
+ :return: Prediction model of the pipeline (or None)
314
+ """
315
+ return self.predictor
316
+
317
+ def add_transform(self, instance: Step) -> None:
318
+ """Add transform Step to the Pipeline
319
+
320
+ :param Step instance: Step to add (must implement transform).
321
+ :raise ValueError: Step must implement transform method.
322
+ """
323
+ if instance and hasattr(instance, 'transform'):
324
+ self._invalidate_fingerprint_cache()
325
+ self._preprocessing_steps.append((str(instance), instance))
326
+ else:
327
+ raise ValueError("Step must implement transform method")
328
+
329
+ def add_resample(self, instance: Step) -> None:
330
+ """Add resample Step to the Pipeline
331
+
332
+ :param Step instance: Step to add (must implement resample).
333
+ :raise ValueError: Step must implement resample method.
334
+ """
335
+ if instance and hasattr(instance, 'resample'):
336
+ self._invalidate_fingerprint_cache()
337
+ self._preprocessing_steps.append((str(instance), instance))
338
+ else:
339
+ raise ValueError("Step must implement resample method")
340
+
341
+ def set_model(self, instance: Step) -> None:
342
+ """
343
+ Set the predict model (Step) of the Pipeline
344
+
345
+ :param Step instance: Step to add (must implement predict).
346
+ :raise ValueError: Step must implement predict method.
347
+ """
348
+ self._invalidate_fingerprint_cache()
349
+ self.predictor = (str(instance), instance)
350
+
351
+ def copy(self) -> IAMLPipeline:
352
+ """Return a copied IAMLPipeline
353
+
354
+ :return: Copied IAMLPipeline instance.
355
+ """
356
+ return deepcopy(self)
357
+
358
+ def pickle(self) -> bytes:
359
+ """Serialize IAMLPipeline to bytes.
360
+ Can be save into a file and reload with pickle.
361
+
362
+ :return: Serialized IAMLPipeline.
363
+ """
364
+ return pickle.dumps(self)
365
+
366
+ @property
367
+ def have_model(self) -> bool:
368
+ """Does the IAMLPipeline have a model set?
369
+
370
+ :return: True if a model has been set.
371
+ """
372
+ return bool(self.predictor)
373
+
374
+ def transform(self, X: pd.DataFrame) -> pd.DataFrame: # pylint: disable=arguments-differ
375
+ """Apply transformers without predict
376
+
377
+ :param pd.DataFrame X: candidate data.
378
+ :return: Transformed DF.
379
+ """
380
+ for _, step in self.transformers:
381
+ X = step.transform(X)
382
+
383
+ return X
384
+
385
+ def predict(self, X: pd.DataFrame, model_only: bool = False, **kwargs: dict) -> list:
386
+ """Run all the steps to predict labels from candidate data
387
+
388
+ :param pd.DataFrame X: Features used as candidate of the pipeline.
389
+ :param bool, optional model_only: True to execute only the model with already transformed
390
+ data. Defaults to False.
391
+ :param dict, optional \\**kwargs: Additional parameters.
392
+ :raise ValueError: Model must have been set before call predict.
393
+ :return: Predicted values
394
+ """
395
+ if not self.have_model:
396
+ raise ValueError("Model need to be set before predict")
397
+
398
+ if not model_only:
399
+ return super().predict(X, **kwargs)
400
+
401
+ if self._trained_columns and isinstance(X, pd.DataFrame):
402
+ X = X.reindex(columns=self._trained_columns, fill_value=0)
403
+ return self.predictor[1].predict(X)
404
+
405
+ def predict_survival_function(
406
+ self,
407
+ X: pd.DataFrame,
408
+ model_only: bool = False,
409
+ **kwargs) -> list:
410
+ """Run all the steps to predict survival function from candidate data
411
+
412
+ :param pd.DataFrame X: Features used as candidate of the pipeline.
413
+ :param bool, optional model_only: True to execute only the model with already transformed
414
+ data. Defaults to False.
415
+ :param dict, optional \\**kwargs: Additional parameters.
416
+ :raise ValueError: Model must have been set before call predict.
417
+ :return: Predicted values
418
+ """
419
+ if not self.have_model:
420
+ raise ValueError("Model need to be set before predict")
421
+
422
+ if not model_only:
423
+ X = self.transform(X, **kwargs)
424
+
425
+ if model_only and self._trained_columns and isinstance(X, pd.DataFrame):
426
+ X = X.reindex(columns=self._trained_columns, fill_value=0)
427
+ return self.predictor[1].predict_survival_function(X)
428
+
429
+ def predict_proba(self, X: pd.DataFrame, model_only: bool = False, **kwargs) -> list:
430
+ """Run all the steps to predict labels from candidate data
431
+
432
+ :param pd.DataFrame X: Features used as candidate of the pipeline.
433
+ :param bool, optional model_only: True to execute only the model with already transformed
434
+ data. Defaults to False.
435
+ :param dict, optional \\**kwargs: Additional parameters.
436
+ :raise ValueError: Model must have been set before call predict.
437
+ :return: Predicted values
438
+ """
439
+ if not self.have_model:
440
+ raise ValueError("Model need to be set before predict")
441
+
442
+ if not model_only:
443
+ return super().predict_proba(X, **kwargs)
444
+
445
+ if self._trained_columns and isinstance(X, pd.DataFrame):
446
+ X = X.reindex(columns=self._trained_columns, fill_value=0)
447
+ return self.predictor[1].predict_proba(X)
448
+
449
+ def __getattribute__(self, attr: str) -> bool:
450
+ """Overload getattr to allow accurate hasattr on predict_proba
451
+
452
+ :param str attr: Attribute to test.
453
+ :raise AttributeError: predict_proba not implemented in this model.
454
+ :return: Does attribute is implemented.
455
+ """
456
+ if attr == 'predict_proba' \
457
+ and not( \
458
+ self.have_model and hasattr(self.predictor[1], 'predict_proba') \
459
+ ):
460
+ raise AttributeError("predict_proba not implemented in this model")
461
+
462
+ return super().__getattribute__(attr)
463
+
464
+ @property
465
+ def optimizable_step(self) -> list['Step']:
466
+ """Return steps whose parameters or choice of implementation can change.
467
+
468
+ :return: List of optimizable step
469
+ """
470
+ return [step for _, step in self.training_steps
471
+ if step.optimizable or step.is_interchangeable]
472
+
473
+ def __eq__(self, other: 'IAMLPipeline') -> bool:
474
+ """Compare two pipelines
475
+
476
+ :param IAMLPipeline other: The pipeline to compare.
477
+ :return: Equal or not ?
478
+ """
479
+ if isinstance(other, IAMLPipeline):
480
+ return self.fingerprint() == other.fingerprint()
481
+ return NotImplemented
482
+
483
+ @property
484
+ def name(self) -> str:
485
+ """Return pipeline formatted name
486
+
487
+ :return: formatted name
488
+ """
489
+ return ' '.join(x.title() for x in str(self.model[0]).split('_'))
490
+
491
+ def explain_model(self, X: pd.DataFrame, nsamples: int = 20) -> Explanation:
492
+ """Explains the model by computing SHAP values on the fitted model.
493
+ Uses the train set as the masker, and the provided set as
494
+ prediction.
495
+
496
+ :param pd.DataFrame X: Prediction set to compute SHAP values for.
497
+ :param int, optional nsamples: Number of samples to pick from the masker to pick feature
498
+ data from for each row in the provided prediction dataset. More samples means more
499
+ accurate SHAP values and longer computing times. Defaults to 20.
500
+ :raise RuntimeError: There is no model to explain.
501
+ :return: Model explanation, with an overview of the most important features, and graphs.
502
+ """
503
+ if not self.have_model:
504
+ raise RuntimeError('There is no model to explain.')
505
+
506
+ def p(pred_data):
507
+ df = pd.DataFrame(pred_data, columns=X.columns)
508
+
509
+ if hasattr(self, 'predict_proba'):
510
+ return self.predict_proba(df)[:, 1]
511
+
512
+ # when the regressor does not implement predict_proba
513
+ return self.predict(df)
514
+
515
+ mask_dataset = self.original_dataset if self.original_dataset is not None \
516
+ and not self.original_dataset.empty else X
517
+
518
+ explainer = shap.KernelExplainer(p, mask_dataset)
519
+ shap_values = explainer.shap_values(X, nsamples=nsamples)
520
+
521
+ shap_explanation = shap.Explanation(
522
+ shap_values,
523
+ base_values=np.tile(explainer.expected_value, (shap_values.shape[0], 1)),
524
+ data=X.to_numpy(),
525
+ feature_names=X.columns.to_list(),
526
+ output_names=X.columns.to_list())
527
+
528
+ return Explanation(shap_explanation)
529
+
530
+ # Implement scikit-learn estimator's methods
531
+ def __sklearn_is_fitted__(self):
532
+ return self.have_model
533
+
534
+ def __sklearn_clone__(self):
535
+ return deepcopy(self)
536
+
537
+ def target_type_(self) -> str:
538
+ """Mimic Scikit-learn API
539
+ Return models target type
540
+ """
541
+ return self.original_dataset.type_of_target
542
+
543
+ # Fingerprint (used by cache)
544
+ def fingerprint(self) -> str:
545
+ """Return a md5 hash that can by use to compare Pipelines
546
+
547
+ :return: md5 sting
548
+ """
549
+ current_version = tuple(
550
+ (id(step), getattr(step, "_config_version", None))
551
+ for _, step in self.training_steps
552
+ )
553
+ if self._fingerprint_cache is None or self._fingerprint_cache_version != current_version:
554
+ to_hash = "\n".join([step.fingerprint() for _, step in self.training_steps])
555
+ self._fingerprint_cache = md5(to_hash.encode()).hexdigest()
556
+ self._fingerprint_cache_version = current_version
557
+ return self._fingerprint_cache
558
+
559
+ def transformers_resamplers_fingerprint(self) -> str:
560
+ """Fingerprint for transformers/resamplers only (used by Candidate)."""
561
+ current_version = tuple(
562
+ (id(step), getattr(step, "_config_version", None))
563
+ for _, step in self._preprocessing_steps
564
+ )
565
+ if self._transform_fingerprint_cache is None \
566
+ or self._transform_fingerprint_cache_version != current_version:
567
+ to_hash = "\n".join([
568
+ step.fingerprint()
569
+ for _, step in self._preprocessing_steps
570
+ ])
571
+ self._transform_fingerprint_cache = md5(to_hash.encode()).hexdigest()
572
+ self._transform_fingerprint_cache_version = current_version
573
+ return self._transform_fingerprint_cache
574
+
575
+ def _invalidate_fingerprint_cache(self) -> None:
576
+ self._fingerprint_cache = None
577
+ self._transform_fingerprint_cache = None
578
+ self._fingerprint_cache_version = None
579
+ self._transform_fingerprint_cache_version = None
580
+
581
+ def bibliography(self, structured: bool) -> str | list[dict]:
582
+ """Return a string listing all step's references or a structured list of dict.
583
+
584
+ :param bool structured: JSON structured bibliography or not.
585
+ :return str | list[dict]: Bibliography.
586
+ """
587
+ references = [reference for step in self.steps
588
+ for reference in step[1].references] \
589
+ + [reference for metric in self.metrics for reference in metric.get_refs()] \
590
+ + [Reference({
591
+ 'year': 2017,
592
+ 'name': 'A Unified Approach to Interpreting Model Predictions',
593
+ 'authors': [
594
+ 'Scott M. Lundberg', 'Su-In Lee'
595
+ ],
596
+ 'doi': 'https://doi.org/10.48550/arXiv.1705.07874',
597
+ 'publisher': 'arXiv preprint arXiv:1705.07874'
598
+ }, 'Shap')]
599
+
600
+ return Reference.bibliography(references, structured)