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,96 @@
1
+ """[PLOT] ROC-AUC Plot"""
2
+ from __future__ import annotations
3
+
4
+ import textwrap
5
+ import io
6
+ from typing import TYPE_CHECKING
7
+ import pandas as pd
8
+ import matplotlib.pyplot as plt
9
+ from sklearn.metrics import roc_curve, auc
10
+
11
+ from ..metric_plot import MetricPlot, capture
12
+ if TYPE_CHECKING:
13
+ from ..iaml_pipeline import IAMLPipeline
14
+
15
+
16
+ class ROCAUCPlot(MetricPlot):
17
+ """[PLOT] ROC-AUC Plot"""
18
+
19
+ title: str = "Receiver Operating Characteristic - Area Under the Curve"
20
+ description: str = textwrap.dedent("""
21
+ The ROC-AUC (Receiver Operating Characteristic - Area Under the Curve) plot is a widely used
22
+ tool to assess the performance of a classification model, especially in the healthcare domain.
23
+ It provides a graphical representation of the model's ability to distinguish between classes,
24
+ such as diagnosing the presence or absence of a medical condition.
25
+
26
+ The ROC curve shows the trade-off between the true positive rate (sensitivity) and false positive
27
+ rate for different threshold values. The AUC score (Area Under the Curve) summarizes the performance
28
+ into a single number, where a value of 1 indicates perfect classification, and 0.5 represents a
29
+ model no better than random guessing.
30
+
31
+ For instance, in a medical setting, you might have a model predicting whether a patient has a
32
+ certain disease or is healthy. The ROC curve would help evaluate how well the model can separate
33
+ patients with the disease from those without. The closer the curve is to the top-left corner and
34
+ the higher the AUC score, the better the model is at distinguishing between the conditions.
35
+
36
+ **Micro-average** and **macro-average** ROC curves are useful when dealing with multiclass classification
37
+ problems (where there are more than two classes).
38
+
39
+ - **Micro-average** ROC aggregates the contributions of all classes and calculates metrics globally
40
+ by counting the total true positives, false positives, true negatives, and false negatives.
41
+ It provides a single ROC curve by combining all classes, treating each decision as a binary one
42
+ (one-vs-rest). This is useful when you care about the overall performance of the classifier across
43
+ all categories.
44
+
45
+ - **Macro-average** ROC computes the ROC curve for each class separately and then averages the results.
46
+ This gives equal weight to all classes, regardless of the number of samples. Macro-average is useful
47
+ when you want to evaluate the model's performance on each class individually, giving each class the
48
+ same importance, regardless of how often it appears in the dataset.
49
+
50
+ Doctors and data scientists use this visualization to ensure that the model performs well across
51
+ different threshold values, making it a critical tool in situations where misdiagnosis could have
52
+ serious consequences.
53
+ """)
54
+
55
+ @capture
56
+ def compute(
57
+ self,
58
+ estimator: IAMLPipeline,
59
+ X: pd.DataFrame,
60
+ y: pd.Series,
61
+ X_train: pd.DataFrame = None,
62
+ y_train: pd.Series = None,
63
+ **kwargs) -> MetricPlot:
64
+ self._binary_image = io.BytesIO()
65
+
66
+ pos_label = None
67
+ if y.dtype not in ['int', 'bool']:
68
+ pos_label = y.iloc[0] if isinstance(y, pd.Series) else y[0]
69
+
70
+ # Predict probabilities
71
+ y_prob = estimator.predict_proba(X)[:, 1]
72
+
73
+ # Compute ROC curve and AUC
74
+ fpr, tpr, _ = roc_curve(y, y_prob, pos_label=pos_label)
75
+ roc_auc = auc(fpr, tpr)
76
+
77
+ # Create the ROC plot
78
+ plt.figure()
79
+ plt.plot(fpr, tpr, color='blue', lw=2, label=f'ROC curve (AUC = {roc_auc:.2f})')
80
+ plt.plot([0, 1], [0, 1], color='grey', lw=2, linestyle='--', label='Random guess')
81
+ plt.xlim([0.0, 1.0])
82
+ plt.ylim([0.0, 1.05])
83
+ plt.xlabel('False Positive Rate')
84
+ plt.ylabel('True Positive Rate')
85
+ plt.title('Receiver Operating Characteristic')
86
+ plt.legend(loc='lower right')
87
+ plt.grid(True)
88
+
89
+ # Save plot to binary image
90
+ plt.savefig(self._binary_image, format='png')
91
+
92
+ return self
93
+
94
+ @classmethod
95
+ def suitable(cls, type_of_target: str) -> bool:
96
+ return type_of_target == 'binary'
@@ -0,0 +1,187 @@
1
+ """[PLOT] Wrap Shap Plot """
2
+ import textwrap
3
+ import io
4
+ from typing import Any
5
+ import matplotlib.pyplot as plt
6
+ import numpy as np
7
+ import shap
8
+
9
+ from ..plot import Plot
10
+
11
+
12
+ class ShapPlot(Plot):
13
+ """[PLOT] Wrap Shap Plot
14
+
15
+ :param str plot_key: The kind of shap plot to use.
16
+ :param shap.Explanation | shap.Cohorts | dict[shap.Explanation] shaps_values: Values used
17
+ by the shap library to compute data.
18
+ :param optional \\*args: Additional parameters.
19
+ :param slice, optional ps: Used to slice shaps_values. Default is None.
20
+ :param optional scatter_feature: Used to select specific shaps_values. Default is None.
21
+ :param optional \\**kwargs: Additional parameters.
22
+
23
+ """
24
+ def __init__(
25
+ self,
26
+ plot_key: str,
27
+ shaps_values: shap.Explanation | shap.Cohorts | dict[shap.Explanation],
28
+ *args,
29
+ ps: slice = None,
30
+ scatter_feature: Any | shap.Cohorts | None = None,
31
+ **kwargs) -> None:
32
+ self.key: str = plot_key
33
+ """The kind of shap plot to perform"""
34
+
35
+ if ps is None:
36
+ ps = slice(0, len(shaps_values))
37
+
38
+ if scatter_feature is None and plot_key == 'scatter':
39
+ scatter_feature = shaps_values.feature_names[0]
40
+
41
+ method, self.title, self.description = self.__plots_informations(plot_key, shaps_values)
42
+ self._binary_image = io.BytesIO()
43
+
44
+ if plot_key == 'scatter':
45
+ method(shaps_values[ps, scatter_feature], *args, show=False, **kwargs)
46
+ elif plot_key == 'force':
47
+ method(shaps_values[ps.start], *args, show=False, matplotlib=True, **kwargs)
48
+ elif plot_key == 'waterfall':
49
+ method(shaps_values[ps.start], *args, show=False, **kwargs)
50
+ else:
51
+ method(shaps_values[ps], *args, show=False, **kwargs)
52
+
53
+ plt.savefig(self._binary_image, bbox_inches='tight')
54
+ plt.close()
55
+
56
+ @classmethod
57
+ def all(
58
+ cls,
59
+ shaps_values: shap.Explanation | shap.Cohorts | dict[shap.Explanation],
60
+ *args,
61
+ **kwargs) -> list['ShapPlot']:
62
+ """Get instance for all kind of Shap Plot
63
+
64
+ :param shap.Explanation | shap.Cohorts | dict[shap.Explanation] shaps_values: Values used
65
+ by the shap library to compute data.
66
+ :param optional \\*args: Additional parameters.
67
+ :param optional \\**kwargs: Additional parameters.
68
+
69
+ :return: list of computed ShapPlot.
70
+ """
71
+ plots = []
72
+ for key in ['force', 'waterfall', 'beeswarm', 'scatter', 'heatmap', 'bar']:
73
+ plots.append(cls(key, shaps_values, *args, **kwargs))
74
+
75
+ return plots
76
+
77
+ def __plots_informations(
78
+ self,
79
+ key: str,
80
+ shap_values: shap.Explanation | shap.Cohorts | dict[shap.Explanation]) -> tuple[str]:
81
+ """
82
+ Returns the title and description for various SHAP plot types in simple terms.
83
+
84
+ This function provides easy-to-understand explanations for SHAP plots, using
85
+ examples from the medical field, to help non-experts interpret how machine learning
86
+ models make predictions.
87
+
88
+ :param str key: The type of SHAP plot
89
+
90
+ * force
91
+ * waterfall
92
+ * beeswarm
93
+ * scatter
94
+ * heatmap
95
+ * bar
96
+ :param shap.Explanation | shap.Cohorts | dict[shap.Explanation] shap_values: Values used
97
+ by the shap library to compute data.
98
+ :return: A title and description of the SHAP plot type, explaining what it shows and how it
99
+ relates to model predictions.
100
+ """
101
+ features = shap_values.feature_names
102
+ values = shap_values[0].values
103
+
104
+ match key:
105
+ case 'force':
106
+ force_shap, force_feature = max(zip(values, features), key=lambda v: abs(v[0]))
107
+ return (shap.plots.force,
108
+ textwrap.dedent("""\
109
+ SHAP Force Plot: Visualizing How Individual Factors Contribute to a Prediction
110
+ """),
111
+ textwrap.dedent(f"""\
112
+ The force plot shows how different factors (e.g., age, cholesterol level, blood
113
+ pressure) push the model’s prediction for an individual patient. It explains
114
+ whether each factor increases or decreases the likelihood of a certain outcome,
115
+ such as a diagnosis of heart disease. Red arrows indicate factors increasing risk,
116
+ while blue arrows show those reducing risk. This plot helps interpret the specific
117
+ impact of each factor for a given prediction.
118
+
119
+ Reading: For this prediction, `{force_feature}` impacts the final prediction
120
+ value by {force_shap:.3f}.
121
+ """))
122
+ case 'waterfall':
123
+ force_shap, force_feature = max(zip(values, features), key=lambda v: abs(v[0]))
124
+ return (shap.plots.waterfall,
125
+ textwrap.dedent("""\
126
+ SHAP Waterfall Plot: Decomposing a Prediction into Its Components
127
+ """),
128
+ textwrap.dedent(f"""\
129
+ The waterfall plot breaks down how each factor influences a single patient's
130
+ prediction by showing the cumulative effect of each factor. Starting from the
131
+ average prediction, it steps through each factor (e.g., age, medication history,
132
+ lab results) to show how the final prediction is reached. This helps in
133
+ understanding the main contributors to a prediction, such as a high blood
134
+ sugar level increasing the risk of diabetes.
135
+
136
+ Reading: For this prediction, `{force_feature}` impacts the final prediction
137
+ value by {force_shap:.3f}.
138
+ """))
139
+ case 'beeswarm':
140
+ return (shap.plots.beeswarm,
141
+ textwrap.dedent("""\
142
+ SHAP Beeswarm Plot: Identifying the Most Important Factors Across All Patients
143
+ """),
144
+ textwrap.dedent("""\
145
+ The beeswarm plot highlights which factors are most important across all patients.
146
+ Each dot represents a patient, with dots positioned based on the factor's impact
147
+ on the prediction (e.g., positive or negative impact on disease risk). For instance,
148
+ a cluster of red dots could show that high blood pressure consistently increases
149
+ heart disease risk. This plot helps find patterns and common trends in the data.
150
+ """))
151
+ case 'scatter':
152
+ return (shap.plots.scatter,
153
+ textwrap.dedent("""\
154
+ SHAP Scatter Plot: Visualizing the Relationship Between a Factor and Prediction
155
+ """),
156
+ textwrap.dedent("""\
157
+ The scatter plot shows the relationship between a specific factor (e.g., body mass index)
158
+ and its SHAP value, which tells us how much it affects the model’s prediction.
159
+ By plotting multiple patients, this plot reveals how changes in a factor (like increasing
160
+ BMI) can lead to higher or lower risk predictions (such as for heart disease).
161
+ """))
162
+ case 'heatmap':
163
+ return (shap.plots.heatmap,
164
+ "SHAP Heatmap: Understanding Factor Importance Across Multiple Patients",
165
+ textwrap.dedent("""\
166
+ The heatmap shows the impact of different factors for many patients, with color
167
+ intensity representing how strongly a factor influences the model's prediction.
168
+ For example, dark red may highlight that high cholesterol is a strong positive
169
+ predictor for heart disease in several patients. This plot helps you see which
170
+ factors are the most influential overall.
171
+ """))
172
+ case 'bar':
173
+ mean_shap = np.abs(shap_values.values).mean(axis=0)
174
+ bar_shap, bar_feature = max(zip(mean_shap, features), key=lambda v: v[0])
175
+ return (shap.plots.bar,
176
+ "SHAP Bar Plot: Ranking the Most Important Factors",
177
+ textwrap.dedent(f"""\
178
+ The bar plot ranks the factors by their overall importance in the model’s
179
+ predictions. Each bar represents a factor (e.g., age, smoking status,
180
+ cholesterol level) and shows how much it contributed to the model’s
181
+ decision-making process across all patients. This helps identify the key
182
+ factors driving predictions, such as high blood pressure being the most
183
+ influential predictor of heart disease.
184
+
185
+ Reading: `{bar_feature}` has an absolute impact of {bar_shap:.3f}
186
+ on the average final prediction value.
187
+ """))
@@ -0,0 +1,241 @@
1
+ """[PLOT] Target distribution 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 ..dataset import Dataset
13
+ from ..plot import StatisticPlot, capture
14
+
15
+
16
+ def _is_missing(value: Any) -> bool:
17
+ if value is None:
18
+ return True
19
+ if isinstance(value, float) and pd.isna(value):
20
+ return True
21
+ return False
22
+
23
+
24
+ def _plot_placeholder(message: str) -> None:
25
+ plt.figure()
26
+ plt.text(0.5, 0.5, message, ha='center', va='center')
27
+ plt.axis('off')
28
+
29
+
30
+ def _extract_target_column(dataframe: pd.DataFrame) -> str | None:
31
+ if dataframe.empty:
32
+ return None
33
+ for col in ('target', 'y', 'target_all', 'y_all'):
34
+ if col in dataframe.columns:
35
+ return col
36
+ if len(dataframe.columns) == 1:
37
+ return dataframe.columns[0]
38
+ return None
39
+
40
+
41
+ def _extract_value_counts(value: Any) -> dict[str, float] | None:
42
+ if _is_missing(value):
43
+ return None
44
+ if isinstance(value, dict):
45
+ if any(key in value for key in ('counts', 'bins', 'bin_edges', 'hist', 'values')):
46
+ return None
47
+ counts: dict[str, float] = {}
48
+ for key, count in value.items():
49
+ if _is_missing(count):
50
+ continue
51
+ counts[str(key)] = float(count)
52
+ return counts or None
53
+ if isinstance(value, (list, tuple, np.ndarray, pd.Series)):
54
+ try:
55
+ items = list(value)
56
+ except TypeError:
57
+ return None
58
+ if not items:
59
+ return None
60
+ if all(isinstance(item, (list, tuple)) and len(item) == 2 for item in items):
61
+ counts = {}
62
+ for key, count in items:
63
+ if _is_missing(count):
64
+ continue
65
+ counts[str(key)] = float(count)
66
+ return counts or None
67
+ return None
68
+
69
+
70
+ def _histogram_from_values(values: np.ndarray) -> tuple[np.ndarray, np.ndarray] | None:
71
+ if values.size == 0:
72
+ return None
73
+ try:
74
+ numeric = values.astype(float)
75
+ except (TypeError, ValueError):
76
+ return None
77
+ numeric = numeric[np.isfinite(numeric)]
78
+ if numeric.size == 0:
79
+ return None
80
+ bins = int(np.sqrt(numeric.size))
81
+ bins = max(5, min(20, bins))
82
+ counts, bin_edges = np.histogram(numeric, bins=bins)
83
+ return counts, bin_edges
84
+
85
+
86
+ def _extract_hist_data(value: Any) -> tuple[np.ndarray, np.ndarray] | None:
87
+ if _is_missing(value):
88
+ return None
89
+ if isinstance(value, dict):
90
+ if 'counts' in value and ('bins' in value or 'bin_edges' in value):
91
+ counts = np.asarray(value['counts'])
92
+ bins = np.asarray(value.get('bins', value.get('bin_edges')))
93
+ return counts, bins
94
+ if 'hist' in value and ('bins' in value or 'bin_edges' in value):
95
+ counts = np.asarray(value['hist'])
96
+ bins = np.asarray(value.get('bins', value.get('bin_edges')))
97
+ return counts, bins
98
+ if 'values' in value:
99
+ return _histogram_from_values(np.asarray(value['values']))
100
+ if isinstance(value, (list, tuple, np.ndarray, pd.Series)):
101
+ if isinstance(value, (list, tuple)) and len(value) == 2:
102
+ first = np.asarray(value[0])
103
+ second = np.asarray(value[1])
104
+ if first.ndim == 1 and second.ndim == 1:
105
+ if first.size == second.size + 1:
106
+ return second, first
107
+ if second.size == first.size + 1:
108
+ return first, second
109
+ if first.size == second.size and first.size > 0:
110
+ bins = np.arange(first.size + 1)
111
+ return first, bins
112
+ return _histogram_from_values(np.asarray(value))
113
+ return None
114
+
115
+
116
+ def _plot_class_counts(counts: dict[str, float], title: str) -> None:
117
+ plt.figure()
118
+ labels = list(counts.keys())
119
+ values = np.asarray(list(counts.values()), dtype=float)
120
+ x = np.arange(len(labels))
121
+ plt.bar(x, values, color='tab:blue')
122
+ plt.xticks(x, labels, rotation=30, ha='right')
123
+ plt.ylabel('Count')
124
+ plt.title(title)
125
+ plt.tight_layout()
126
+
127
+
128
+ def _plot_histogram(counts: np.ndarray, bins: np.ndarray, title: str) -> None:
129
+ if bins.size != counts.size + 1:
130
+ _plot_placeholder("Invalid histogram data")
131
+ return
132
+ plt.figure()
133
+ widths = np.diff(bins)
134
+ plt.bar(bins[:-1], counts, width=widths, align='edge', edgecolor='black')
135
+ plt.xlabel('Target')
136
+ plt.ylabel('Count')
137
+ plt.title(title)
138
+ plt.tight_layout()
139
+
140
+
141
+ class TargetDistributionPlot(StatisticPlot):
142
+ """[PLOT] Target Distribution Plot."""
143
+
144
+ name: str = "Target Distribution"
145
+ _description: str = textwrap.dedent("""\
146
+ Target distribution plots summarize the target values.
147
+ """)
148
+ _description_long: str = textwrap.dedent("""\
149
+ This plot shows the distribution of the target variable. For regression tasks,
150
+ it renders a histogram of the target values. For classification tasks, it shows
151
+ counts per class label.
152
+ """)
153
+ refs: list[dict] = []
154
+
155
+ title: str = "Target distribution"
156
+ description: str = textwrap.dedent("""\
157
+ The target distribution plot shows the distribution of the target values.
158
+ """)
159
+ group_by_feature: bool = False
160
+
161
+ def __str__(self) -> str:
162
+ return 'target_distribution'
163
+
164
+ @capture
165
+ def compute(
166
+ self,
167
+ dataframe: pd.DataFrame,
168
+ dataset: Dataset | None = None,
169
+ base_name: str | None = None,
170
+ **kwargs,
171
+ ) -> 'TargetDistributionPlot':
172
+ """Compute the target distribution plot."""
173
+ self._binary_image = io.BytesIO()
174
+
175
+ title = "Target distribution"
176
+ if base_name:
177
+ title = f"Target distribution: {base_name}"
178
+
179
+ if not dataframe.empty:
180
+ target_col = _extract_target_column(dataframe)
181
+ row = None
182
+ row_key = str(self)
183
+ if row_key in dataframe.index:
184
+ row = dataframe.loc[row_key]
185
+ elif 'value_counts' in dataframe.index:
186
+ row = dataframe.loc['value_counts']
187
+ elif 'histogram' in dataframe.index:
188
+ row = dataframe.loc['histogram']
189
+
190
+ if isinstance(row, pd.DataFrame):
191
+ row = row.iloc[0] if not row.empty else None
192
+
193
+ if row is not None and target_col is not None:
194
+ value = row.get(target_col)
195
+ counts = _extract_value_counts(value)
196
+ if counts:
197
+ _plot_class_counts(counts, title)
198
+ plt.savefig(self._binary_image, format='png')
199
+ return self
200
+
201
+ hist = _extract_hist_data(value)
202
+ if hist is not None:
203
+ counts_arr, bins = hist
204
+ _plot_histogram(
205
+ np.asarray(counts_arr, dtype=float),
206
+ np.asarray(bins, dtype=float),
207
+ title,
208
+ )
209
+ plt.savefig(self._binary_image, format='png')
210
+ return self
211
+
212
+ if dataset is None:
213
+ _plot_placeholder("Target distribution not available")
214
+ plt.savefig(self._binary_image, format='png')
215
+ return self
216
+
217
+ if dataset.type_of_target == 'survival':
218
+ _plot_placeholder("Target distribution not available for survival targets")
219
+ plt.savefig(self._binary_image, format='png')
220
+ return self
221
+
222
+ y_values = pd.Series(dataset.y)
223
+ if dataset.type_of_target == 'continuous':
224
+ numeric = pd.to_numeric(y_values, errors='coerce').to_numpy()
225
+ numeric = numeric[np.isfinite(numeric)]
226
+ hist = _histogram_from_values(numeric)
227
+ if hist is None:
228
+ _plot_placeholder("No numeric target values available")
229
+ else:
230
+ counts_arr, bins = hist
231
+ _plot_histogram(counts_arr, bins, title)
232
+ else:
233
+ counts_series = y_values.value_counts(dropna=False)
234
+ if counts_series.empty:
235
+ _plot_placeholder("No target labels available")
236
+ else:
237
+ counts = {str(key): float(val) for key, val in counts_series.items()}
238
+ _plot_class_counts(counts, title)
239
+
240
+ plt.savefig(self._binary_image, format='png')
241
+ return self