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
@@ -0,0 +1,101 @@
1
+
2
+ """
3
+ Base class of IAML Optimizer. Optimizer receive a pool of Candidates,
4
+ optimize parameters and return a new pool of candidate
5
+ """
6
+ from copy import deepcopy
7
+ import random
8
+ import time
9
+ from ..candidate import Candidate
10
+ from .optimizer import Optimizer
11
+ from ..step import Step
12
+ from ..logger import Logger
13
+
14
+ class RandomOptimizer(Optimizer):
15
+ def __init__(self, duration:int=None, max_iterations=50):
16
+ """
17
+ Initialize the Random Search Optimizer.
18
+ :param candidates: List of Candidate objects to optimize.
19
+ :param max_iterations: Number of random samples to evaluate.
20
+ :param patience: Number of iterations without improvement before stopping.
21
+ """
22
+ self.max_iterations = max_iterations
23
+ self.current_iteration = 0
24
+ self.initial_modifier:float = 5
25
+ self.max_candidates = 40
26
+ self.first_candidates = None
27
+
28
+
29
+ def _randomize_hyperparameters(self, candidate):
30
+ """Randomly modifies the hyperparameters of a given candidate."""
31
+ for _, step in candidate.pipeline.steps:
32
+ for key, config in step.configuration.items():
33
+ # Get random values
34
+
35
+ if type(config['value']) in [int, float]: # Numeric value ? Let's apply multiplier
36
+ is_int = isinstance(config['value'], int)
37
+
38
+ new_value = None
39
+ if 'range' in config: # Random in range
40
+ new_value = random.uniform(*config['range'])
41
+ else: # Strong multiplier -> kind of random
42
+ # Randomly choose a positive or negative editing
43
+ if bool(random.getrandbits(1)):
44
+ # Negative -> Multiply value by something between 0.01 and 1
45
+ change_rate = random.uniform(0.01, 1)
46
+ new_value = config['value']*change_rate
47
+ else:
48
+ # Positive -> Multiply value by something between 1
49
+ # and the max modificator in configuration
50
+ change_rate = random.uniform(1, self.initial_modifier)
51
+ new_value = config['value']*change_rate
52
+
53
+ # Value was a int ? Round it to keep it int
54
+ if is_int:
55
+ new_value = round(new_value)
56
+
57
+ if not self.__valide_config(config, new_value):
58
+ # Cancel is the new value is not correct.
59
+ new_value = config['value']
60
+
61
+ # Categorical value, choose randomly one of them
62
+ elif 'categorical' in config.keys():
63
+ new_value = random.choice(config['categorical'])
64
+ elif isinstance(config['value'], bool):
65
+ # Bool value, choose randomly beetwen True and False
66
+ new_value = random.choice([True, False])
67
+ else: # Other value ? Just keep it
68
+ new_value = config['value']
69
+
70
+ step.configure(key, new_value)
71
+
72
+ return candidate
73
+
74
+ def run(self, candidates):
75
+ """
76
+ Perform a single iteration of random search optimization.
77
+ :param evaluate_fn: Function that takes a Candidate object and returns a performance score.
78
+ """
79
+ candidates = candidates[0:self.max_candidates//2]
80
+ new_candidates = []
81
+ for candidate in candidates:
82
+ new_candidates.append(self._randomize_hyperparameters(deepcopy(candidate)))
83
+
84
+ self.current_iteration += 1
85
+ return candidates + new_candidates
86
+
87
+ @property
88
+ def finished(self) -> bool:
89
+ """
90
+ Does optimisation is finished ?
91
+
92
+ Returns:
93
+ bool: finished ?
94
+ """
95
+ return self.current_iteration >= self.max_iterations
96
+
97
+ # Does the configuration is valid or not ?
98
+ def __valide_config(self, config:dict, value:any) -> bool:
99
+ if 'range' not in config.keys():
100
+ return True
101
+ return config['range'][0] <= value <= config['range'][1]
iaml/plot.py ADDED
@@ -0,0 +1,138 @@
1
+ """[PLOT] Parent of all others Plot, implement the default behavior"""
2
+ from __future__ import annotations
3
+
4
+ import base64
5
+ from functools import wraps
6
+ from typing import Any, TYPE_CHECKING
7
+
8
+ import io
9
+ import matplotlib.pyplot as plt
10
+ import pandas as pd
11
+
12
+ if TYPE_CHECKING:
13
+ from .iaml_pipeline import IAMLPipeline
14
+
15
+
16
+ def capture(func) -> Any:
17
+ """A decorator to capture a matplotlib plot into a BytesIO object and return it
18
+ as binary data instead of showing it.
19
+
20
+ :return: binary plot data.
21
+ """
22
+ @wraps(func)
23
+ def wrapper(*args, **kwargs):
24
+ ret = func(*args, **kwargs)
25
+ plt.close()
26
+ return ret
27
+
28
+ return wrapper
29
+
30
+
31
+ class Plot:
32
+ """[PLOT] Parent of all others Plot, implement the default behavior"""
33
+
34
+ title: str = "Here is the plot title"
35
+ """Plot title"""
36
+
37
+ description: str = "Here is an explanation of how this plot work"
38
+ """Plot description"""
39
+
40
+ def __init__(self, *_):
41
+ self._binary_image: io.BytesIO = None
42
+ """The generated plot image"""
43
+
44
+ @property
45
+ def image(self) -> bytes:
46
+ """Get binary representation of the plot image
47
+
48
+ :raise AttributeError: Plot must be computed before.
49
+ :return: The image bytes.
50
+ """
51
+ if self._binary_image is not None:
52
+ return self._binary_image.getvalue()
53
+
54
+ raise AttributeError("Plot must be computed before")
55
+
56
+ @property
57
+ def b64_image(self) -> str:
58
+ """Get b64 representation of the plot image
59
+
60
+ :return: b64 string image.
61
+ """
62
+ return base64.b64encode(self.image).decode()
63
+
64
+ def compute(
65
+ self,
66
+ estimator: IAMLPipeline,
67
+ X: pd.DataFrame,
68
+ y: pd.Series,
69
+ **kwargs) -> None:
70
+ """Compute plot given X, y.
71
+ Must be overloaded by children classes
72
+
73
+ :param IAMLPipeline estimator: The pipeline we compute the plot on.
74
+ :param pd.DataFrame X: The dataset we wanna compute plot on.
75
+ :param pd.Series y: The dataset target we wanna compute plot on.
76
+ :param optional \\**kwargs: Additional parameters for plotting.
77
+ """
78
+ raise NotImplementedError('Subclass must implement abstract method')
79
+
80
+ @classmethod
81
+ def suitable(cls, type_of_target: str) -> bool: # pylint: disable=unused-argument
82
+ """
83
+ Evaluates whether the plot is relevant to the type of target to
84
+ predict.
85
+
86
+ :param str type_of_target: Type of target to predict.
87
+ :return: Whether the plot is relevant.
88
+ """
89
+ return False
90
+
91
+ def to_json(self, data_format: str='binary') -> dict[str, Any]:
92
+ """Convert plot into json with name, description and b64 image
93
+
94
+ :param str data_format: Type of data we want. Can be 'binary' or 'b64'.
95
+ :raise AttributeError: Invalid data format.
96
+ :return: A dictonnary containing plot informations and data.
97
+ """
98
+ match data_format:
99
+ case 'binary':
100
+ data = self.image
101
+ case 'b64':
102
+ data = self.b64_image
103
+ case _:
104
+ raise AttributeError('Invalid data format')
105
+
106
+ return {'title': self.title,
107
+ 'description': self.description,
108
+ 'image': data}
109
+
110
+ def to_markdown(self) -> str:
111
+ """Return plot as markdown format
112
+
113
+ :return: Markdown formatted plot.
114
+ """
115
+ return "\n\n".join([
116
+ f"# {self.title}",
117
+ self.description,
118
+ f"![{self.title}](data:image/png;base64,{self.b64_image})"
119
+ ])
120
+
121
+
122
+ class StatisticPlot(Plot):
123
+ """Base class for descriptive statistics plots."""
124
+
125
+ enabled: bool = True
126
+ """Whether this plot should be considered for rendering."""
127
+
128
+ group_by_feature: bool = False
129
+ """Whether to build one plot per feature column."""
130
+
131
+ def compute(self, dataframe: pd.DataFrame, **kwargs) -> 'StatisticPlot':
132
+ """Compute plot given a descriptive statistics dataframe.
133
+
134
+ :param pd.DataFrame dataframe: The descriptive statistics dataframe.
135
+ :param optional \\**kwargs: Additional parameters for plotting.
136
+ :return: A StatisticPlot object.
137
+ """
138
+ raise NotImplementedError('Subclass must implement abstract method')
iaml/plots/__init__.py ADDED
@@ -0,0 +1,32 @@
1
+ """All plots"""
2
+
3
+ # Classifier
4
+ from .class_prediction_error_plot import ClassPredictionErrorPlot
5
+ from .classification_report_plot import ClassificationReportPlot
6
+ from .confusion_matrix_plot import ConfusionMatrixPlot
7
+ from .rocauc_plot import ROCAUCPlot
8
+ from .precision_recall_curve_plot import PrecisionRecallCurvePlot
9
+
10
+ # Regressor
11
+ from .residual_plot import ResidualsPlot
12
+ from .prediction_error_plot import PredictionErrorPlot
13
+
14
+ # Survival
15
+ from .kaplan_meier_comparison_plot import KaplanMeierModelComparisonPlot
16
+ from .cumulative_hazard_plot import CumulativeHazardModelComparisonPlot
17
+ from .roc_dynamique_curve_plot import ROCDynamiqueCurvePlot
18
+ from .shap_plot import ShapPlot
19
+
20
+ # Descriptive statistics
21
+ from .bar_plot import BarPlot
22
+ from .line_plot import LinePlot
23
+ from .histogram_plot import HistogramPlot
24
+ from .box_plot import BoxPlot
25
+ from .violin_plot import ViolinPlot
26
+ from .density_plot import DensityPlot
27
+ from .qq_plot import QQPlot
28
+ from .correlation_heatmap_plot import CorrelationHeatmapPlot
29
+ from .missingness_heatmap_plot import MissingnessHeatmapPlot
30
+ from .pair_plot import PairPlot
31
+ from .target_distribution_plot import TargetDistributionPlot
32
+ from .outlier_plot import OutlierPlot
iaml/plots/bar_plot.py ADDED
@@ -0,0 +1,141 @@
1
+ """[PLOT] Bar plot for descriptive statistics."""
2
+ from __future__ import annotations
3
+
4
+ import io
5
+ import textwrap
6
+ import numpy as np
7
+ import pandas as pd
8
+ import matplotlib.pyplot as plt
9
+
10
+ from ..plot import StatisticPlot, capture
11
+
12
+
13
+ def _is_missing(value: object) -> bool:
14
+ if value is None:
15
+ return True
16
+ if isinstance(value, float) and pd.isna(value):
17
+ return True
18
+ return False
19
+
20
+
21
+ def _infer_base_name(columns: list[str]) -> str:
22
+ for col in columns:
23
+ if isinstance(col, str) and col.endswith('_all'):
24
+ return col[:-4]
25
+ first = columns[0] if columns else ''
26
+ first_str = str(first)
27
+ return first_str.rsplit('_', 1)[0] if '_' in first_str else first_str
28
+
29
+
30
+ def _column_label(column: str, base_name: str | None) -> str:
31
+ column_str = str(column)
32
+ base_name_str = str(base_name) if base_name is not None else None
33
+ if base_name_str:
34
+ if column_str == base_name_str:
35
+ return 'all'
36
+ prefix = f"{base_name_str}_"
37
+ if column_str.startswith(prefix):
38
+ return column_str[len(prefix):]
39
+ if column_str.endswith('_all'):
40
+ return 'all'
41
+ return column_str.rsplit('_', 1)[-1] if '_' in column_str else 'all'
42
+
43
+
44
+ def _plot_placeholder(message: str) -> None:
45
+ plt.figure()
46
+ plt.text(0.5, 0.5, message, ha='center', va='center')
47
+ plt.axis('off')
48
+
49
+
50
+ class BarPlot(StatisticPlot):
51
+ """[PLOT] Bar Plot."""
52
+
53
+ title: str = "Bar plot"
54
+ description: str = textwrap.dedent("""\
55
+ The bar plot shows descriptive statistics for categorical or numerical columns.
56
+ Categorical plots show value counts per class, while numerical plots show metrics
57
+ such as mean or variance per class.
58
+ """)
59
+ group_by_feature: bool = True
60
+
61
+ @capture
62
+ def compute(
63
+ self,
64
+ dataframe: pd.DataFrame,
65
+ base_name: str | None = None,
66
+ **kwargs) -> 'BarPlot':
67
+ """Compute bar plot statistics."""
68
+ self._binary_image = io.BytesIO()
69
+
70
+ if dataframe.empty:
71
+ _plot_placeholder("No statistics available")
72
+ plt.savefig(self._binary_image, format='png')
73
+ return self
74
+
75
+ base_name = base_name or _infer_base_name(list(dataframe.columns))
76
+
77
+ if 'value_counts' in dataframe.index and dataframe.loc['value_counts'].notna().any():
78
+ categories = []
79
+ seen = set()
80
+ for col in dataframe.columns:
81
+ values = dataframe.at['value_counts', col]
82
+ if _is_missing(values):
83
+ continue
84
+ for cat, _ in values:
85
+ if cat not in seen:
86
+ seen.add(cat)
87
+ categories.append(cat)
88
+
89
+ if not categories:
90
+ _plot_placeholder("No categorical statistics available")
91
+ plt.savefig(self._binary_image, format='png')
92
+ return self
93
+
94
+ labels = [_column_label(col, base_name) for col in dataframe.columns]
95
+ data = {cat: [] for cat in categories}
96
+ for col in dataframe.columns:
97
+ values = dataframe.at['value_counts', col]
98
+ value_dict = dict(values) if not _is_missing(values) else {}
99
+ for cat in categories:
100
+ data[cat].append(value_dict.get(cat, 0))
101
+ df = pd.DataFrame(data, index=labels)
102
+
103
+ plt.figure()
104
+ x = np.arange(len(df.index))
105
+ width = min(0.8 / max(len(df.columns), 1), 0.2)
106
+ for i, category in enumerate(df.columns):
107
+ offset = (i - (len(df.columns) - 1) / 2) * width
108
+ plt.bar(x + offset, df[category], width, label=category)
109
+ plt.xticks(x, labels)
110
+ plt.legend()
111
+ plt.ylabel('Count')
112
+ plt.title(f'[CAT] {base_name} Statistics')
113
+ plt.tight_layout()
114
+ elif 'mean' in dataframe.index and dataframe.loc['mean'].notna().any():
115
+ numeric_df = dataframe.copy()
116
+ for row in ['mode', 'value_counts', 'null_count', 'count']:
117
+ if row in numeric_df.index:
118
+ numeric_df = numeric_df.drop(index=row)
119
+ numeric_df = numeric_df.dropna(how='all')
120
+ if numeric_df.empty:
121
+ _plot_placeholder("No numeric statistics available")
122
+ plt.savefig(self._binary_image, format='png')
123
+ return self
124
+
125
+ labels = [_column_label(col, base_name) for col in numeric_df.columns]
126
+ x = np.arange(len(numeric_df.index))
127
+ width = min(0.8 / max(len(numeric_df.columns), 1), 0.25)
128
+ plt.figure()
129
+ for i, col in enumerate(numeric_df.columns):
130
+ offset = (i - (len(numeric_df.columns) - 1) / 2) * width
131
+ plt.bar(x + offset, numeric_df[col], width, label=labels[i])
132
+ plt.xticks(x, numeric_df.index, rotation=45, ha='right')
133
+ plt.legend()
134
+ plt.ylabel('Value')
135
+ plt.title(f'[NUM] {base_name} Statistics')
136
+ plt.tight_layout()
137
+ else:
138
+ _plot_placeholder("No statistics available for bar plot")
139
+
140
+ plt.savefig(self._binary_image, format='png')
141
+ return self
iaml/plots/box_plot.py ADDED
@@ -0,0 +1,166 @@
1
+ """[PLOT] Box plot for descriptive statistics."""
2
+ from __future__ import annotations
3
+
4
+ import io
5
+ import textwrap
6
+ from typing import Any
7
+
8
+ import numpy as np
9
+ import pandas as pd
10
+ import matplotlib.pyplot as plt
11
+
12
+ from ..data_type import DataType
13
+ from ..dataset import Dataset
14
+ from ..plot import StatisticPlot, capture
15
+
16
+
17
+ def _is_missing(value: Any) -> bool:
18
+ if value is None:
19
+ return True
20
+ if isinstance(value, float) and pd.isna(value):
21
+ return True
22
+ return False
23
+
24
+
25
+ def _infer_base_name(columns: list[str]) -> str:
26
+ for col in columns:
27
+ if isinstance(col, str) and col.endswith('_all'):
28
+ return col[:-4]
29
+ first = columns[0] if columns else ''
30
+ first_str = str(first)
31
+ return first_str.rsplit('_', 1)[0] if '_' in first_str else first_str
32
+
33
+
34
+ def _column_label(column: str, base_name: str | None) -> str:
35
+ column_str = str(column)
36
+ base_name_str = str(base_name) if base_name is not None else None
37
+ if base_name_str:
38
+ if column_str == base_name_str:
39
+ return 'all'
40
+ prefix = f"{base_name_str}_"
41
+ if column_str.startswith(prefix):
42
+ return column_str[len(prefix):]
43
+ if column_str.endswith('_all'):
44
+ return 'all'
45
+ return column_str.rsplit('_', 1)[-1] if '_' in column_str else column_str
46
+
47
+
48
+ def _base_from_column(column: str) -> str:
49
+ column_str = str(column)
50
+ if column_str.endswith('_all'):
51
+ return column_str[:-4]
52
+ return column_str.rsplit('_', 1)[0] if '_' in column_str else column_str
53
+
54
+
55
+ def _plot_placeholder(message: str) -> None:
56
+ plt.figure()
57
+ plt.text(0.5, 0.5, message, ha='center', va='center')
58
+ plt.axis('off')
59
+
60
+
61
+ def _extract_box_stats(dataframe: pd.DataFrame, column: str) -> dict[str, float] | None:
62
+ required_rows = {
63
+ 'min': 'whislo',
64
+ 'max': 'whishi',
65
+ 'quantile_0.25': 'q1',
66
+ 'quantile_0.5': 'med',
67
+ 'quantile_0.75': 'q3',
68
+ }
69
+ stats: dict[str, float] = {}
70
+ for row, key in required_rows.items():
71
+ if row not in dataframe.index:
72
+ return None
73
+ value = dataframe.at[row, column]
74
+ if _is_missing(value):
75
+ return None
76
+ try:
77
+ value = float(value)
78
+ except (TypeError, ValueError):
79
+ return None
80
+ if not np.isfinite(value):
81
+ return None
82
+ stats[key] = value
83
+ stats['fliers'] = []
84
+ return stats
85
+
86
+
87
+ class BoxPlot(StatisticPlot):
88
+ """[PLOT] Box Plot."""
89
+
90
+ name: str = "Box Plot"
91
+ _description: str = textwrap.dedent("""\
92
+ Box plots summarize numeric distributions per class.
93
+ """)
94
+ _description_long: str = textwrap.dedent("""\
95
+ Box plots visualize min, quartiles, and max values for numeric columns
96
+ per class, using precomputed descriptive statistics.
97
+ """)
98
+ refs: list[dict] = []
99
+
100
+ title: str = "Box plot"
101
+ description: str = textwrap.dedent("""\
102
+ The box plot shows numerical distributions per class.
103
+ """)
104
+ group_by_feature: bool = True
105
+
106
+ def __str__(self) -> str:
107
+ return 'boxplot'
108
+
109
+ @capture
110
+ def compute(
111
+ self,
112
+ dataframe: pd.DataFrame,
113
+ base_name: str | None = None,
114
+ dataset: Dataset | None = None,
115
+ **kwargs) -> 'BoxPlot':
116
+ """Compute box plot statistics."""
117
+ self._binary_image = io.BytesIO()
118
+
119
+ if dataframe.empty:
120
+ _plot_placeholder("No statistics available")
121
+ plt.savefig(self._binary_image, format='png')
122
+ return self
123
+
124
+ base_name = base_name or _infer_base_name(list(dataframe.columns))
125
+
126
+ if dataset is not None:
127
+ numeric_columns = set(dataset.get_columns_names_by_type(DataType.NUMERIC))
128
+ if base_name in numeric_columns:
129
+ columns_to_show = [
130
+ col for col in dataframe.columns
131
+ if _base_from_column(col) == base_name
132
+ ]
133
+ else:
134
+ columns_to_show = [
135
+ col for col in dataframe.columns
136
+ if _base_from_column(col) in numeric_columns
137
+ ]
138
+ else:
139
+ columns_to_show = list(dataframe.columns)
140
+
141
+ if not columns_to_show:
142
+ _plot_placeholder("No numeric statistics available")
143
+ plt.savefig(self._binary_image, format='png')
144
+ return self
145
+
146
+ entries: list[dict[str, Any]] = []
147
+ for col in columns_to_show:
148
+ stats = _extract_box_stats(dataframe, col)
149
+ if stats is None:
150
+ continue
151
+ stats['label'] = _column_label(col, base_name)
152
+ entries.append(stats)
153
+
154
+ if not entries:
155
+ _plot_placeholder("No box plot statistics available")
156
+ plt.savefig(self._binary_image, format='png')
157
+ return self
158
+
159
+ plt.figure(figsize=(max(4.0, 0.9 * len(entries)), 4.0))
160
+ plt.bxp(entries, showfliers=False)
161
+ plt.ylabel('Value')
162
+ if base_name:
163
+ plt.title(f"Box plot: {base_name}")
164
+ plt.tight_layout()
165
+ plt.savefig(self._binary_image, format='png')
166
+ return self
@@ -0,0 +1,37 @@
1
+ """[PLOT] Class Prediction Error Plot"""
2
+ import textwrap
3
+
4
+ from yellowbrick.classifier import ClassPredictionError
5
+
6
+ from ..metric_plot import MetricPlot, yellowbrick_plot
7
+
8
+
9
+ @yellowbrick_plot(ClassPredictionError)
10
+ class ClassPredictionErrorPlot(MetricPlot):
11
+ """[PLOT] Class Prediction Error Plot"""
12
+
13
+ title: str = "Prediction Error Plot"
14
+ description: str = textwrap.dedent("""
15
+ The Class Prediction Error is a visualization that helps understand how well a machine
16
+ learning model is performing in predicting medical conditions or diagnoses. It shows
17
+ both the correct predictions made by the model and the mistakes it makes for each condition.
18
+
19
+ For example, imagine you have a model that’s trained to identify different diseases
20
+ from patient data, such as predicting whether someone has diabetes, hypertension,
21
+ or is healthy. The Class Prediction Error plot would show, for each of these conditions,
22
+ how many times the model correctly identified the disease and how many times it made a
23
+ wrong prediction.
24
+
25
+ For instance, if the model predicts "diabetes" for a patient who actually has "hypertension,"
26
+ the plot will highlight this error. Similarly, it will also show how often the model
27
+ correctly identifies "healthy" patients versus when it mistakenly predicts they have
28
+ a disease.
29
+
30
+ This visualization is especially helpful for doctors and data scientists because it clearly
31
+ shows where the model is making errors, making it easier to improve its accuracy, which is
32
+ critical in healthcare where correct predictions can have a big impact on patient outcomes.
33
+ """)
34
+
35
+ @classmethod
36
+ def suitable(cls, type_of_target: str) -> bool:
37
+ return type_of_target in ['binary', 'multiclass', 'multilabel-indicator']
@@ -0,0 +1,35 @@
1
+ """[PLOT] Classification Report Plot"""
2
+ import textwrap
3
+
4
+ from yellowbrick.classifier import ClassificationReport
5
+
6
+ from ..metric_plot import MetricPlot, yellowbrick_plot
7
+
8
+
9
+ @yellowbrick_plot(ClassificationReport)
10
+ class ClassificationReportPlot(MetricPlot):
11
+ """[PLOT] Classification Report Plot"""
12
+
13
+ title: str = "Classification report"
14
+ description: str = textwrap.dedent("""
15
+ The Classification Report is a visual tool to evaluate the performance of a machine learning model
16
+ on classification tasks, such as diagnosing medical conditions. This plot provides key metrics for
17
+ each class (like diseases or health conditions) that the model is trained to identify.
18
+
19
+ It includes metrics such as precision, recall, F1-score, and support for each class. These metrics
20
+ are essential to understanding how well the model is identifying true positives (correct diagnoses),
21
+ minimizing false positives (incorrect diagnoses), and balancing between precision and recall.
22
+
23
+ For example, if you have a model classifying conditions like 'healthy', 'diabetes', and 'hypertension',
24
+ the Classification Report plot will show you the precision (how many of the predicted conditions were correct),
25
+ recall (how many actual conditions were correctly identified), and the F1-score (the harmonic mean of precision
26
+ and recall). This is especially important in healthcare to ensure the model provides balanced and accurate results
27
+ across all classes, improving both diagnosis and patient outcomes.
28
+
29
+ Doctors and data scientists use this visualization to easily compare the model's performance on different conditions,
30
+ aiding in model refinement and ensuring robust diagnostic predictions.
31
+ """)
32
+
33
+ @classmethod
34
+ def suitable(cls, type_of_target: str) -> bool:
35
+ return type_of_target in ['binary', 'multiclass', 'multilabel-indicator']
@@ -0,0 +1,34 @@
1
+ """[PLOT] Confusion Matrix Plot"""
2
+ import textwrap
3
+
4
+ from yellowbrick.classifier import ConfusionMatrix
5
+
6
+ from ..metric_plot import MetricPlot, yellowbrick_plot
7
+
8
+
9
+ @yellowbrick_plot(ConfusionMatrix)
10
+ class ConfusionMatrixPlot(MetricPlot):
11
+ """[PLOT] Confusion Matrix Plot"""
12
+
13
+ title: str = "Confusion Matrix"
14
+ description: str = textwrap.dedent("""
15
+ The Confusion Matrix is a powerful visualization tool used to assess how well a classification model
16
+ is performing, particularly in identifying different medical conditions or diagnostic categories.
17
+ It provides a clear breakdown of true positive, false positive, true negative, and false negative rates.
18
+
19
+ For example, in a healthcare setting, if a model is trained to classify whether a patient has 'diabetes',
20
+ 'hypertension', or is 'healthy', the Confusion Matrix will show how many times the model made correct predictions
21
+ (true positives and true negatives) and where it made mistakes (false positives and false negatives).
22
+
23
+ Each row in the matrix represents the actual condition of the patient, and each column represents the predicted
24
+ condition. This makes it easy to spot patterns in the model's predictions, such as whether it tends to misclassify
25
+ one condition as another.
26
+
27
+ Doctors and data scientists rely on this plot to understand not only how often the model is right but also the
28
+ types of errors it makes. This information is crucial in healthcare, where reducing misdiagnoses can significantly
29
+ improve patient outcomes.
30
+ """)
31
+
32
+ @classmethod
33
+ def suitable(cls, type_of_target: str) -> bool:
34
+ return type_of_target in ['binary', 'multiclass', 'multilabel-indicator']