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,228 @@
1
+ """[PLOT] Pair 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 _plot_placeholder(message: str) -> None:
18
+ plt.figure()
19
+ plt.text(0.5, 0.5, message, ha='center', va='center')
20
+ plt.axis('off')
21
+
22
+
23
+ def _sample_frame(frame: pd.DataFrame, max_points: int) -> pd.DataFrame:
24
+ if frame.shape[0] <= max_points:
25
+ return frame
26
+ indices = np.linspace(0, frame.shape[0] - 1, max_points, dtype=int)
27
+ return frame.iloc[indices]
28
+
29
+
30
+ def _to_numeric_frame(frame: pd.DataFrame) -> pd.DataFrame:
31
+ return frame.apply(pd.to_numeric, errors='coerce')
32
+
33
+
34
+ def _extract_row_values(row: pd.Series) -> pd.DataFrame | None:
35
+ if not isinstance(row, pd.Series):
36
+ return None
37
+ data: dict[str, np.ndarray] = {}
38
+ min_len: int | None = None
39
+ for col, value in row.items():
40
+ if value is None or (isinstance(value, float) and pd.isna(value)):
41
+ continue
42
+ arr: np.ndarray | None = None
43
+ if isinstance(value, dict) and 'values' in value:
44
+ arr = np.asarray(value['values'])
45
+ elif isinstance(value, (list, tuple, np.ndarray, pd.Series)):
46
+ arr = np.asarray(value)
47
+ if arr is None:
48
+ continue
49
+ arr = arr.ravel()
50
+ if arr.size == 0:
51
+ continue
52
+ if min_len is None or arr.size < min_len:
53
+ min_len = int(arr.size)
54
+ data[str(col)] = arr
55
+
56
+ if not data or min_len is None or min_len < 2:
57
+ return None
58
+
59
+ trimmed = {col: values[:min_len] for col, values in data.items()}
60
+ return pd.DataFrame(trimmed)
61
+
62
+
63
+ def _frame_from_stats(dataframe: pd.DataFrame, key: str) -> pd.DataFrame | None:
64
+ if dataframe.empty or key not in dataframe.index:
65
+ return None
66
+ row = dataframe.loc[key]
67
+ if isinstance(row, pd.DataFrame):
68
+ if row.empty:
69
+ return None
70
+ row = row.iloc[0]
71
+ return _extract_row_values(row)
72
+
73
+
74
+ def _looks_like_raw_data(dataframe: pd.DataFrame) -> bool:
75
+ if dataframe.shape[0] < 2 or dataframe.shape[1] < 2:
76
+ return False
77
+ stats_keys = {
78
+ 'mean', 'median', 'min', 'max', 'std', 'variance', 'count', 'mode',
79
+ 'null_count', 'missing_rate', 'value_counts', 'range', 'iqr',
80
+ }
81
+ if any(str(label) in stats_keys for label in dataframe.index):
82
+ return False
83
+ if dataframe.index.dtype == object and dataframe.index.nunique() <= 2:
84
+ return False
85
+ numeric = _to_numeric_frame(dataframe)
86
+ return np.isfinite(numeric.to_numpy()).any()
87
+
88
+
89
+ def _select_numeric_frame(
90
+ dataframe: pd.DataFrame,
91
+ dataset: Dataset | None,
92
+ max_features: int,
93
+ max_points: int,
94
+ ) -> pd.DataFrame | None:
95
+ if dataset is not None:
96
+ numeric_columns = dataset.get_columns_names_by_type(DataType.NUMERIC)
97
+ if not numeric_columns:
98
+ return None
99
+ columns = numeric_columns[:max_features]
100
+ frame = dataset.X[columns]
101
+ frame = _to_numeric_frame(frame)
102
+ frame = _sample_frame(frame, max_points)
103
+ return frame
104
+
105
+ if _looks_like_raw_data(dataframe):
106
+ frame = _to_numeric_frame(dataframe)
107
+ if frame.shape[1] > max_features:
108
+ frame = frame.iloc[:, :max_features]
109
+ frame = _sample_frame(frame, max_points)
110
+ return frame
111
+
112
+ return None
113
+
114
+
115
+ def _figure_size(n_features: int) -> tuple[float, float]:
116
+ size = float(min(12.0, max(4.0, 2.2 * n_features)))
117
+ return size, size
118
+
119
+
120
+ class PairPlot(StatisticPlot):
121
+ """[PLOT] Pair Plot."""
122
+
123
+ name: str = "Pair Plot"
124
+ _description: str = textwrap.dedent("""\
125
+ Pair plots show pairwise scatter plots between numeric features.
126
+ """)
127
+ _description_long: str = textwrap.dedent("""\
128
+ This plot renders a scatter matrix for a small number of numeric features,
129
+ highlighting pairwise relationships and marginal distributions.
130
+ """)
131
+ refs: list[dict] = []
132
+
133
+ title: str = "Pair plot"
134
+ description: str = textwrap.dedent("""\
135
+ The pair plot shows pairwise scatter plots for small numeric feature sets.
136
+ """)
137
+ group_by_feature: bool = False
138
+
139
+ def __str__(self) -> str:
140
+ return 'pairplot'
141
+
142
+ @capture
143
+ def compute(
144
+ self,
145
+ dataframe: pd.DataFrame,
146
+ dataset: Dataset | None = None,
147
+ base_name: str | None = None,
148
+ **kwargs,
149
+ ) -> 'PairPlot':
150
+ """Compute pair plot statistics."""
151
+ self._binary_image = io.BytesIO()
152
+
153
+ max_features = int(kwargs.get('max_features', 6))
154
+ max_points = int(kwargs.get('max_points', 800))
155
+ max_features = max(2, max_features)
156
+ max_points = max(50, max_points)
157
+
158
+ frame = _frame_from_stats(dataframe, str(self))
159
+ if frame is None:
160
+ for candidate in ('pair_plot', 'scatter_matrix', 'scatter'):
161
+ frame = _frame_from_stats(dataframe, candidate)
162
+ if frame is not None:
163
+ break
164
+
165
+ if frame is None:
166
+ frame = _select_numeric_frame(dataframe, dataset, max_features, max_points)
167
+
168
+ if frame is None or frame.empty:
169
+ _plot_placeholder("Pair plot data not available")
170
+ plt.savefig(self._binary_image, format='png')
171
+ return self
172
+
173
+ frame = _to_numeric_frame(frame)
174
+ frame = frame.dropna(how='all')
175
+ frame = frame.loc[:, frame.notna().any(axis=0)]
176
+ if frame.shape[1] < 2 or frame.shape[0] < 2:
177
+ _plot_placeholder("Not enough numeric data for pair plot")
178
+ plt.savefig(self._binary_image, format='png')
179
+ return self
180
+
181
+ if frame.shape[1] > max_features:
182
+ frame = frame.iloc[:, :max_features]
183
+
184
+ columns = [str(col) for col in frame.columns]
185
+ n_features = len(columns)
186
+ fig_w, fig_h = _figure_size(n_features)
187
+ fig, axes = plt.subplots(n_features, n_features, figsize=(fig_w, fig_h))
188
+
189
+ values = frame.to_numpy(dtype=float)
190
+ for i in range(n_features):
191
+ for j in range(n_features):
192
+ ax = axes[i, j]
193
+ if i == j:
194
+ data = values[:, i]
195
+ data = data[np.isfinite(data)]
196
+ if data.size == 0:
197
+ ax.text(0.5, 0.5, "No data", ha='center', va='center')
198
+ ax.axis('off')
199
+ else:
200
+ bins = int(np.sqrt(data.size))
201
+ bins = max(5, min(15, bins))
202
+ ax.hist(data, bins=bins, color='tab:blue', alpha=0.7)
203
+ else:
204
+ x = values[:, j]
205
+ y = values[:, i]
206
+ mask = np.isfinite(x) & np.isfinite(y)
207
+ if not np.any(mask):
208
+ ax.text(0.5, 0.5, "No data", ha='center', va='center')
209
+ ax.axis('off')
210
+ else:
211
+ ax.scatter(x[mask], y[mask], s=10, alpha=0.6, color='tab:blue')
212
+
213
+ if i == n_features - 1:
214
+ ax.set_xlabel(columns[j], rotation=45, ha='right')
215
+ else:
216
+ ax.set_xticklabels([])
217
+ if j == 0:
218
+ ax.set_ylabel(columns[i])
219
+ else:
220
+ ax.set_yticklabels([])
221
+
222
+ title = "Pair plot"
223
+ if base_name:
224
+ title = f"Pair plot: {base_name}"
225
+ fig.suptitle(title)
226
+ plt.tight_layout(rect=(0, 0, 1, 0.95))
227
+ plt.savefig(self._binary_image, format='png')
228
+ return self
@@ -0,0 +1,86 @@
1
+ """
2
+ [PLOT] Precision-Recall Curve Plot
3
+ """
4
+ from __future__ import annotations
5
+
6
+ import textwrap
7
+ import io
8
+ from typing import TYPE_CHECKING
9
+ import pandas as pd
10
+ import matplotlib.pyplot as plt
11
+ from sklearn.metrics import precision_recall_curve, average_precision_score
12
+
13
+ from ..metric_plot import MetricPlot, capture
14
+ if TYPE_CHECKING:
15
+ from ..iaml_pipeline import IAMLPipeline
16
+
17
+
18
+ class PrecisionRecallCurvePlot(MetricPlot):
19
+ """[PLOT] Precision-Recall Curve Plot"""
20
+
21
+ title: str = "Precision-Recall Curve"
22
+ description: str = textwrap.dedent("""
23
+ The Precision-Recall Curve is a valuable visualization tool for evaluating the performance of
24
+ a classification model, particularly in healthcare where identifying the correct balance between
25
+ precision (positive predictive value) and recall (sensitivity or true positive rate) is critical.
26
+
27
+ Precision measures how many of the predicted positive cases (e.g., disease diagnoses) were actually correct,
28
+ while recall indicates how well the model identifies all the true positive cases. The Precision-Recall Curve
29
+ shows this trade-off across different threshold settings of the model.
30
+
31
+ In healthcare, for example, if a model is predicting whether a patient has a disease, the Precision-Recall Curve
32
+ will help you understand how the model performs when prioritizing minimizing false positives (increasing precision)
33
+ versus maximizing the detection of true positives (increasing recall). This is especially important in imbalanced
34
+ datasets where one condition (e.g., healthy patients) dominates over others (e.g., rare diseases).
35
+
36
+ Doctors and data scientists rely on this visualization to optimize the model based on specific healthcare priorities,
37
+ such as minimizing missed diagnoses or reducing unnecessary treatments, making the Precision-Recall Curve an essential
38
+ tool for improving patient outcomes.
39
+ """)
40
+
41
+ @capture
42
+ def compute( # pylint: disable=arguments-differ
43
+ self,
44
+ estimator: IAMLPipeline,
45
+ X: pd.DataFrame,
46
+ y: pd.Series,
47
+ **kwargs) -> MetricPlot:
48
+ """
49
+ Compute plot given X, y.
50
+ """
51
+ self._binary_image = io.BytesIO()
52
+
53
+ pos_label = None
54
+ if y.dtype == 'int':
55
+ pos_label = 1
56
+ elif y.dtype == 'bool':
57
+ pos_label = True
58
+ else:
59
+ pos_label = y.iloc[0] if isinstance(y, pd.Series) else y[0]
60
+
61
+ # Predict probabilities for the positive class
62
+ y_prob = estimator.predict_proba(X)[:, 1]
63
+
64
+ # Compute Precision-Recall curve
65
+ precision, recall, _ = precision_recall_curve(y, y_prob, pos_label=pos_label)
66
+ average_precision = average_precision_score(y, y_prob, pos_label=pos_label)
67
+
68
+ # Create the Precision-Recall plot
69
+ plt.figure()
70
+ plt.plot(recall, precision, color='blue',
71
+ lw=2, label=f'Precision-Recall curve (AP = {average_precision:.2f})')
72
+ plt.xlabel('Recall')
73
+ plt.ylabel('Precision')
74
+ plt.title('Precision-Recall Curve')
75
+ plt.legend(loc='lower left')
76
+ plt.grid(True)
77
+
78
+ # Save plot to binary image
79
+ plt.savefig(self._binary_image, format='png')
80
+ plt.close()
81
+
82
+ return self
83
+
84
+ @classmethod
85
+ def suitable(cls, type_of_target: str) -> bool:
86
+ return type_of_target == 'binary'
@@ -0,0 +1,34 @@
1
+ """[PLOT] Prediction Error Plot"""
2
+ import textwrap
3
+
4
+ from yellowbrick.regressor import PredictionError
5
+
6
+ from ..metric_plot import MetricPlot, yellowbrick_plot
7
+
8
+
9
+ @yellowbrick_plot(PredictionError)
10
+ class PredictionErrorPlot(MetricPlot):
11
+ """[PLOT] Prediction Error Plot"""
12
+
13
+ title: str = "Prediction Error"
14
+ description: str = textwrap.dedent("""
15
+ The Prediction Error Plot is a diagnostic tool used to visualize the performance of regression models.
16
+ It helps assess how well a model's predictions align with the actual values in a continuous prediction
17
+ setting, such as predicting medical measurements like blood pressure, heart rate, or glucose levels.
18
+
19
+ The plot shows the actual target values on the x-axis and the predicted values on the y-axis. A perfect
20
+ model would have all points lying on a 45-degree diagonal line, representing perfect predictions. Deviations
21
+ from this line indicate errors in the predictions.
22
+
23
+ For instance, if you're building a model to predict a patient's blood pressure based on certain health metrics,
24
+ the Prediction Error Plot will show how closely the model's predictions match the actual measurements. If the
25
+ points are scattered away from the diagonal, it indicates that the model is making large prediction errors.
26
+
27
+ This plot is useful for identifying both systematic errors (consistent overestimation or underestimation) and
28
+ random errors (scattering of points) in the model, which is crucial in healthcare settings where accurate
29
+ predictions can directly impact patient care.
30
+ """)
31
+
32
+ @classmethod
33
+ def suitable(cls, type_of_target: str) -> bool:
34
+ return type_of_target == 'continuous'
iaml/plots/qq_plot.py ADDED
@@ -0,0 +1,220 @@
1
+ """[PLOT] QQ 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
+ from scipy import stats
12
+
13
+ from ..data_type import DataType
14
+ from ..dataset import Dataset
15
+ from ..plot import StatisticPlot, capture
16
+
17
+
18
+ def _is_missing(value: Any) -> bool:
19
+ if value is None:
20
+ return True
21
+ if isinstance(value, float) and pd.isna(value):
22
+ return True
23
+ return False
24
+
25
+
26
+ def _column_label(column: str, base_name: str | None) -> str:
27
+ column_str = str(column)
28
+ base_name_str = str(base_name) if base_name is not None else None
29
+ if base_name_str:
30
+ if column_str == base_name_str:
31
+ return 'all'
32
+ prefix = f"{base_name_str}_"
33
+ if column_str.startswith(prefix):
34
+ return column_str[len(prefix):]
35
+ if column_str.endswith('_all'):
36
+ return 'all'
37
+ return column_str
38
+
39
+
40
+ def _plot_placeholder(message: str) -> None:
41
+ plt.figure()
42
+ plt.text(0.5, 0.5, message, ha='center', va='center')
43
+ plt.axis('off')
44
+
45
+
46
+ def _sample_values(values: np.ndarray, max_points: int = 2000) -> np.ndarray:
47
+ if values.size <= max_points:
48
+ return values
49
+ indices = np.linspace(0, values.size - 1, max_points, dtype=int)
50
+ return values[indices]
51
+
52
+
53
+ def _sanitize_pair(
54
+ theoretical: np.ndarray,
55
+ ordered: np.ndarray,
56
+ ) -> tuple[np.ndarray, np.ndarray] | None:
57
+ try:
58
+ theoretical_arr = np.asarray(theoretical, dtype=float).ravel()
59
+ ordered_arr = np.asarray(ordered, dtype=float).ravel()
60
+ except (TypeError, ValueError):
61
+ return None
62
+ if theoretical_arr.size == 0 or ordered_arr.size == 0:
63
+ return None
64
+ if theoretical_arr.size != ordered_arr.size:
65
+ return None
66
+ mask = np.isfinite(theoretical_arr) & np.isfinite(ordered_arr)
67
+ if not np.any(mask):
68
+ return None
69
+ theoretical_arr = theoretical_arr[mask]
70
+ ordered_arr = ordered_arr[mask]
71
+ order = np.argsort(theoretical_arr)
72
+ return theoretical_arr[order], ordered_arr[order]
73
+
74
+
75
+ def _qq_from_sample(values: np.ndarray) -> tuple[np.ndarray, np.ndarray] | None:
76
+ if values.size < 2:
77
+ return None
78
+ try:
79
+ numeric = values.astype(float)
80
+ except (TypeError, ValueError):
81
+ return None
82
+ numeric = numeric[np.isfinite(numeric)]
83
+ if numeric.size < 2:
84
+ return None
85
+ ordered = np.sort(numeric)
86
+ ordered = _sample_values(ordered)
87
+ n = ordered.size
88
+ if n < 2:
89
+ return None
90
+ probs = (np.arange(1, n + 1) - 0.5) / n
91
+ theoretical = stats.norm.ppf(probs)
92
+ return theoretical, ordered
93
+
94
+
95
+ def _extract_qq_data(value: Any) -> tuple[np.ndarray, np.ndarray] | None:
96
+ if _is_missing(value):
97
+ return None
98
+ if isinstance(value, dict):
99
+ if 'theoretical' in value and 'ordered' in value:
100
+ return _sanitize_pair(value['theoretical'], value['ordered'])
101
+ if 'theoretical_quantiles' in value and 'sample_quantiles' in value:
102
+ return _sanitize_pair(value['theoretical_quantiles'], value['sample_quantiles'])
103
+ if 'values' in value:
104
+ return _qq_from_sample(np.asarray(value['values']))
105
+ if 'sample' in value:
106
+ return _qq_from_sample(np.asarray(value['sample']))
107
+ if isinstance(value, (list, tuple, np.ndarray, pd.Series)):
108
+ if isinstance(value, (list, tuple)) and len(value) == 2:
109
+ first = np.asarray(value[0])
110
+ second = np.asarray(value[1])
111
+ if first.ndim == 1 and second.ndim == 1 and first.size == second.size:
112
+ return _sanitize_pair(first, second)
113
+ return _qq_from_sample(np.asarray(value))
114
+ return None
115
+
116
+
117
+ class QQPlot(StatisticPlot):
118
+ """[PLOT] QQ Plot."""
119
+
120
+ name: str = "QQ Plot"
121
+ _description: str = textwrap.dedent("""\
122
+ QQ plots compare numeric distributions to a normal reference.
123
+ """)
124
+ _description_long: str = textwrap.dedent("""\
125
+ This plot compares ordered sample values against theoretical quantiles
126
+ of the normal distribution, highlighting departures from normality.
127
+ """)
128
+ refs: list[dict] = []
129
+
130
+ title: str = "QQ plot"
131
+ description: str = textwrap.dedent("""\
132
+ The QQ plot compares numeric columns to a normal distribution.
133
+ """)
134
+ group_by_feature: bool = True
135
+
136
+ def __str__(self) -> str:
137
+ return 'qqplot'
138
+
139
+ @capture
140
+ def compute(
141
+ self,
142
+ dataframe: pd.DataFrame,
143
+ base_name: str | None = None,
144
+ dataset: Dataset | None = None,
145
+ **kwargs,
146
+ ) -> 'QQPlot':
147
+ """Compute QQ plot statistics."""
148
+ self._binary_image = io.BytesIO()
149
+
150
+ if dataframe.empty:
151
+ _plot_placeholder("No statistics available")
152
+ plt.savefig(self._binary_image, format='png')
153
+ return self
154
+
155
+ qq_key = str(self)
156
+ if qq_key not in dataframe.index:
157
+ for candidate in ('qq', 'qq_plot'):
158
+ if candidate in dataframe.index:
159
+ qq_key = candidate
160
+ break
161
+ else:
162
+ _plot_placeholder("QQ plot statistics not available")
163
+ plt.savefig(self._binary_image, format='png')
164
+ return self
165
+
166
+ if dataset is not None:
167
+ numeric_columns = set(dataset.get_columns_names_by_type(DataType.NUMERIC))
168
+ columns_to_show = [col for col in dataframe.columns if col in numeric_columns]
169
+ else:
170
+ columns_to_show = list(dataframe.columns)
171
+
172
+ qq_row = dataframe.loc[qq_key]
173
+ entries: list[tuple[str, tuple[np.ndarray, np.ndarray]]] = []
174
+ for col in columns_to_show:
175
+ qq_data = _extract_qq_data(qq_row.get(col))
176
+ if qq_data is not None:
177
+ entries.append((col, qq_data))
178
+
179
+ if not entries:
180
+ _plot_placeholder("No numeric QQ statistics available")
181
+ plt.savefig(self._binary_image, format='png')
182
+ return self
183
+
184
+ n_plots = len(entries)
185
+ n_cols = 1 if n_plots == 1 else 2
186
+ n_rows = int(np.ceil(n_plots / n_cols))
187
+ fig, axes = plt.subplots(n_rows, n_cols, figsize=(6 * n_cols, 3.5 * n_rows))
188
+ axes_list = np.atleast_1d(axes).ravel()
189
+
190
+ for ax, (col, (theoretical, ordered)) in zip(axes_list, entries):
191
+ if theoretical.size == 0 or ordered.size == 0 or theoretical.size != ordered.size:
192
+ ax.text(0.5, 0.5, "Invalid QQ data", ha='center', va='center')
193
+ ax.axis('off')
194
+ continue
195
+ ax.scatter(theoretical, ordered, s=12, alpha=0.7, color='tab:blue')
196
+ mean = float(np.mean(ordered))
197
+ std = float(np.std(ordered, ddof=1)) if ordered.size > 1 else 0.0
198
+ x_min = float(np.min(theoretical))
199
+ x_max = float(np.max(theoretical))
200
+ line_x = np.array([x_min, x_max], dtype=float)
201
+ if np.isfinite(std) and std > 0:
202
+ line_y = mean + std * line_x
203
+ else:
204
+ line_y = np.array([mean, mean], dtype=float)
205
+ ax.plot(line_x, line_y, color='red', linewidth=1)
206
+ ax.set_title(_column_label(col, base_name))
207
+ ax.set_xlabel('Theoretical quantiles')
208
+ ax.set_ylabel('Ordered values')
209
+
210
+ for ax in axes_list[len(entries):]:
211
+ ax.axis('off')
212
+
213
+ if base_name:
214
+ fig.suptitle(f"QQ plot: {base_name}")
215
+ plt.tight_layout(rect=(0, 0, 1, 0.95))
216
+ else:
217
+ plt.tight_layout()
218
+
219
+ plt.savefig(self._binary_image, format='png')
220
+ return self
@@ -0,0 +1,38 @@
1
+ """
2
+ [PLOT] Residuals Plot
3
+ """
4
+ import textwrap
5
+
6
+ from yellowbrick.regressor import ResidualsPlot as ybResidualsPlot
7
+
8
+ from ..metric_plot import MetricPlot, yellowbrick_plot
9
+
10
+
11
+ @yellowbrick_plot(ybResidualsPlot)
12
+ class ResidualsPlot(MetricPlot):
13
+ """[PLOT] Residuals Plot"""
14
+
15
+ title: str = "Residuals Plot"
16
+ description: str = textwrap.dedent("""
17
+ The Residuals Plot is a diagnostic tool used to evaluate the performance of a regression model.
18
+ In the context of predicting continuous medical outcomes, such as blood pressure, cholesterol levels,
19
+ or other measurements, this plot helps assess how well the model's predictions match the actual
20
+ observed values.
21
+
22
+ Residuals are the differences between the predicted values and the actual values. A well-performing
23
+ regression model should have residuals that are randomly scattered around zero. Patterns in the residuals
24
+ (such as curvature or clustering) can indicate that the model is not capturing certain relationships in
25
+ the data.
26
+
27
+ For example, if you're building a model to predict a patient's cholesterol level based on various
28
+ health metrics, the Residuals Plot would show whether the model consistently overestimates or underestimates
29
+ values or if there are systematic errors.
30
+
31
+ Doctors and data scientists use this plot to detect whether the model is biased in its predictions and
32
+ whether certain patterns or trends remain unexplained, which can be critical in refining models used
33
+ for predicting medical outcomes.
34
+ """)
35
+
36
+ @classmethod
37
+ def suitable(cls, type_of_target: str) -> bool:
38
+ return type_of_target == 'continuous'
@@ -0,0 +1,79 @@
1
+ """[PLOT] ROC Dynamique Curve for Survival Models using sksurv"""
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 numpy as np
9
+ import matplotlib.pyplot as plt
10
+ from sksurv.metrics import cumulative_dynamic_auc
11
+
12
+ from ..metric_plot import MetricPlot, capture
13
+ if TYPE_CHECKING:
14
+ from ..iaml_pipeline import IAMLPipeline
15
+
16
+
17
+ class ROCDynamiqueCurvePlot(MetricPlot):
18
+ """[PLOT] ROC Dynamique Curve for Survival Models using sksurv"""
19
+
20
+ title: str = "ROC Dynamique Curve"
21
+ description: str = textwrap.dedent("""
22
+ This curve represents how well a predictive survival model is able to distinguish
23
+ between patients who experience an event (like death or a heart attack) at different
24
+ points in time and those who do not. The y-axis shows the AUC (Area Under the Curve),
25
+ which is a measure of how good the model is at making this distinction—the closer to 1,
26
+ the better the model performs. The x-axis represents time, showing different follow-up
27
+ periods after the initial observation.
28
+
29
+ As time progresses, the curve helps us see if the model's predictions remain accurate
30
+ or start to decline. For example, in a medical study predicting patient survival after
31
+ a heart attack, this curve would indicate how well the model distinguishes between
32
+ patients who pass away versus those who survive, over several months or years.
33
+ A high AUC value means the model is very good at predicting outcomes, while a lower
34
+ value suggests it struggles to differentiate between high-risk and low-risk patients as
35
+ time goes on.""")
36
+
37
+ @capture
38
+ def compute(
39
+ self,
40
+ estimator: IAMLPipeline,
41
+ X: pd.DataFrame,
42
+ y: pd.Series,
43
+ X_train: pd.DataFrame = None,
44
+ y_train: pd.Series = None,
45
+ **kwargs) -> MetricPlot:
46
+ self._binary_image = io.BytesIO()
47
+
48
+ # Compute time-dependent ROC AUC for each time point
49
+ _, event_times = zip(*y)
50
+
51
+ max_val = max(event_times) - 0.1 if isinstance(max(event_times), float) \
52
+ else max(event_times)
53
+ times = np.arange(min(event_times), max_val)
54
+
55
+ # Calculate cumulative dynamic AUC (time-dependent ROC AUC)
56
+ aucs, _ = cumulative_dynamic_auc(
57
+ np.array(y_train, dtype=[('event', 'bool'), ('time', 'float')]),
58
+ np.array(y, dtype=[('event', 'bool'), ('time', 'float')]),
59
+ estimator.predict(X),
60
+ times
61
+ )
62
+
63
+ # Plot the time-dependent ROC AUC over time
64
+ plt.plot(times, aucs)
65
+ plt.xlabel("Temps de suivi")
66
+ plt.ylabel("AUC dynamique cumulative")
67
+ plt.title('ROC dynamique curve')
68
+ plt.ylim([0, 1])
69
+ plt.legend()
70
+ plt.grid(True)
71
+
72
+ plt.savefig(self._binary_image, format='png')
73
+ plt.close()
74
+
75
+ return self
76
+
77
+ @classmethod
78
+ def suitable(cls, type_of_target: str) -> bool:
79
+ return type_of_target == 'survival'