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,201 @@
1
+ """[PLOT] Correlation heatmap 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 _plot_placeholder(message: str) -> None:
26
+ plt.figure()
27
+ plt.text(0.5, 0.5, message, ha='center', va='center')
28
+ plt.axis('off')
29
+
30
+
31
+ def _to_numeric_frame(frame: pd.DataFrame) -> pd.DataFrame:
32
+ numeric = frame.apply(pd.to_numeric, errors='coerce')
33
+ return numeric
34
+
35
+
36
+ def _square_corr_from_dataframe(frame: pd.DataFrame) -> pd.DataFrame | None:
37
+ if frame.empty:
38
+ return None
39
+ if frame.shape[0] != frame.shape[1]:
40
+ return None
41
+ if set(frame.index) != set(frame.columns):
42
+ return None
43
+ ordered = frame.reindex(index=frame.index, columns=frame.index)
44
+ numeric = _to_numeric_frame(ordered)
45
+ if not np.isfinite(numeric.to_numpy()).any():
46
+ return None
47
+ return numeric
48
+
49
+
50
+ def _corr_from_value(value: Any) -> pd.DataFrame | None:
51
+ if _is_missing(value):
52
+ return None
53
+ if isinstance(value, pd.DataFrame):
54
+ return _square_corr_from_dataframe(value)
55
+ if isinstance(value, dict):
56
+ matrix = None
57
+ if 'matrix' in value:
58
+ matrix = value['matrix']
59
+ elif 'values' in value:
60
+ matrix = value['values']
61
+ if matrix is None:
62
+ return None
63
+ df = pd.DataFrame(matrix)
64
+ labels = value.get('labels')
65
+ if labels is not None:
66
+ df.index = labels
67
+ df.columns = labels
68
+ if 'index' in value:
69
+ df.index = value['index']
70
+ if 'columns' in value:
71
+ df.columns = value['columns']
72
+ return _square_corr_from_dataframe(df)
73
+ if isinstance(value, (list, tuple, np.ndarray)):
74
+ array = np.asarray(value)
75
+ if array.ndim == 2:
76
+ return _square_corr_from_dataframe(pd.DataFrame(array))
77
+ return None
78
+
79
+
80
+ def _corr_from_row(row: pd.Series) -> pd.DataFrame | None:
81
+ if not isinstance(row, pd.Series):
82
+ return None
83
+ for item in row:
84
+ corr = _corr_from_value(item)
85
+ if corr is not None:
86
+ return corr
87
+ return None
88
+
89
+
90
+ def _extract_corr_dataframe(dataframe: pd.DataFrame) -> pd.DataFrame | None:
91
+ corr_df = _square_corr_from_dataframe(dataframe)
92
+ if corr_df is not None:
93
+ return corr_df
94
+ for key in ('correlation', 'correlation_matrix', 'corr', 'corr_matrix'):
95
+ if key in dataframe.index:
96
+ return _corr_from_row(dataframe.loc[key])
97
+ return None
98
+
99
+
100
+ def _figure_size(n_features: int) -> float:
101
+ if n_features <= 1:
102
+ return 4.0
103
+ return float(min(12.0, max(4.5, 0.5 * n_features + 2.0)))
104
+
105
+
106
+ class CorrelationHeatmapPlot(StatisticPlot):
107
+ """[PLOT] Correlation Heatmap Plot."""
108
+
109
+ name: str = "Correlation Heatmap"
110
+ _description: str = textwrap.dedent("""\
111
+ Correlation heatmaps summarize relationships between numeric features.
112
+ """)
113
+ _description_long: str = textwrap.dedent("""\
114
+ This plot displays a correlation matrix for numeric columns, helping
115
+ identify strong positive or negative relationships.
116
+ """)
117
+ refs: list[dict] = []
118
+
119
+ title: str = "Correlation heatmap"
120
+ description: str = textwrap.dedent("""\
121
+ The correlation heatmap visualizes pairwise correlations among numeric columns.
122
+ """)
123
+ group_by_feature: bool = False
124
+
125
+ def __str__(self) -> str:
126
+ return 'correlation'
127
+
128
+ @capture
129
+ def compute(
130
+ self,
131
+ dataframe: pd.DataFrame,
132
+ dataset: Dataset | None = None,
133
+ base_name: str | None = None,
134
+ **kwargs,
135
+ ) -> 'CorrelationHeatmapPlot':
136
+ """Compute correlation heatmap statistics."""
137
+ self._binary_image = io.BytesIO()
138
+
139
+ if dataframe.empty:
140
+ _plot_placeholder("No statistics available")
141
+ plt.savefig(self._binary_image, format='png')
142
+ return self
143
+
144
+ corr_df = _extract_corr_dataframe(dataframe)
145
+
146
+ if corr_df is None and dataset is not None:
147
+ numeric_columns = dataset.get_columns_names_by_type(DataType.NUMERIC)
148
+ if numeric_columns:
149
+ corr_df = dataset.X[numeric_columns].corr()
150
+
151
+ if corr_df is None or corr_df.empty:
152
+ _plot_placeholder("Correlation statistics not available")
153
+ plt.savefig(self._binary_image, format='png')
154
+ return self
155
+
156
+ corr_df = _to_numeric_frame(corr_df)
157
+ if not np.isfinite(corr_df.to_numpy()).any():
158
+ _plot_placeholder("Correlation statistics not available")
159
+ plt.savefig(self._binary_image, format='png')
160
+ return self
161
+
162
+ if dataset is not None:
163
+ numeric_columns = dataset.get_columns_names_by_type(DataType.NUMERIC)
164
+ if numeric_columns:
165
+ available = [col for col in corr_df.index if col in numeric_columns]
166
+ if available:
167
+ corr_df = corr_df.loc[available, available]
168
+
169
+ if corr_df.empty:
170
+ _plot_placeholder("Correlation statistics not available")
171
+ plt.savefig(self._binary_image, format='png')
172
+ return self
173
+
174
+ n_features = corr_df.shape[0]
175
+ fig_size = _figure_size(n_features)
176
+ fig, ax = plt.subplots(figsize=(fig_size, fig_size))
177
+
178
+ values = corr_df.to_numpy(dtype=float)
179
+ masked = np.ma.masked_invalid(values)
180
+ image = ax.imshow(masked, cmap='coolwarm', vmin=-1, vmax=1)
181
+ plt.colorbar(image, ax=ax, fraction=0.046, pad=0.04)
182
+
183
+ labels = [str(label) for label in corr_df.columns]
184
+ ticks = np.arange(n_features)
185
+ ax.set_xticks(ticks)
186
+ ax.set_yticks(ticks)
187
+ ax.set_xticklabels(labels, rotation=45, ha='right')
188
+ ax.set_yticklabels(labels)
189
+
190
+ label_size = 10 if n_features <= 12 else 8 if n_features <= 20 else 6
191
+ ax.tick_params(axis='both', which='major', labelsize=label_size)
192
+
193
+ title = "Correlation heatmap"
194
+ if base_name:
195
+ title = f"Correlation heatmap: {base_name}"
196
+ ax.set_title(title)
197
+ ax.set_xlabel('Features')
198
+ ax.set_ylabel('Features')
199
+ plt.tight_layout()
200
+ plt.savefig(self._binary_image, format='png')
201
+ return self
@@ -0,0 +1,72 @@
1
+ """[PLOT] Cumulative Hazard Model Comparison Plot using sksurv"""
2
+ import io
3
+ from typing import TYPE_CHECKING
4
+ import textwrap
5
+ from sksurv.nonparametric import nelson_aalen_estimator
6
+ import matplotlib.pyplot as plt
7
+ import numpy as np
8
+ import pandas as pd
9
+ from ..metric_plot import MetricPlot, capture
10
+ if TYPE_CHECKING:
11
+ from ..iaml_pipeline import IAMLPipeline
12
+
13
+
14
+ class CumulativeHazardModelComparisonPlot(MetricPlot):
15
+ """[PLOT] Cumulative Hazard Model Comparison Plot using sksurv"""
16
+
17
+ title: str = "Cumulative Hazard"
18
+ description: str = textwrap.dedent("""
19
+ The Cumulative Hazard Model Comparison Plot is a diagnostic tool used to evaluate the performance of
20
+ survival models by comparing predicted cumulative hazard functions against the observed cumulative hazards.
21
+
22
+ The x-axis represents time, and the y-axis represents the cumulative hazard. Ideally,
23
+ the model-predicted cumulative hazard curves should closely align with the observed curves,
24
+ indicating good model performance. Discrepancies between the two curves highlight
25
+ areas where the model's predictions diverge from reality, signaling potential issues with
26
+ the model's predictive ability.
27
+
28
+ Additionally, a Cox proportional hazards model can be trained to serve as a baseline for comparison.
29
+ """)
30
+
31
+ @capture
32
+ def compute(
33
+ self,
34
+ estimator: 'IAMLPipeline',
35
+ X: pd.DataFrame,
36
+ y: pd.Series,
37
+ X_train: pd.DataFrame = None,
38
+ y_train: pd.Series = None,
39
+ **kwargs) -> MetricPlot:
40
+ self._binary_image = io.BytesIO()
41
+
42
+ # Fit the cumulative hazard model on observed data
43
+ event, time = zip(*y)
44
+
45
+ # Observed cumulative hazard using nelson_aalen_estimator
46
+ time, cumulative_hazard = nelson_aalen_estimator(event, time)
47
+ plt.step(time, cumulative_hazard, where="post", label="Observed", color='blue')
48
+
49
+ # Current model prediction
50
+ hazard_predictions = estimator.predict_cumulative_hazard_function(X)
51
+
52
+ mean_hazard_prob = np.mean([fn.y for fn in hazard_predictions], axis=0)
53
+ mean_hazard_time = hazard_predictions[0].x
54
+
55
+ plt.step(mean_hazard_time, mean_hazard_prob,
56
+ where="post", label="Model prediction", color="green")
57
+
58
+ # Customize and save the plot
59
+ plt.title("Cumulative Hazard Curve vs Model Predicted")
60
+ plt.xlabel("Time")
61
+ plt.ylabel("Cumulative Hazard")
62
+ plt.ylim([0, 1])
63
+ plt.legend()
64
+
65
+ plt.savefig(self._binary_image, format='png')
66
+ plt.close()
67
+
68
+ return self
69
+
70
+ @classmethod
71
+ def suitable(cls, type_of_target: str) -> bool:
72
+ return type_of_target == 'survival'
@@ -0,0 +1,210 @@
1
+ """[PLOT] Density 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 _kde_from_values(values: np.ndarray) -> tuple[np.ndarray, np.ndarray] | None:
54
+ if values.size < 2:
55
+ return None
56
+ try:
57
+ numeric = values.astype(float)
58
+ except (TypeError, ValueError):
59
+ return None
60
+ numeric = numeric[np.isfinite(numeric)]
61
+ if numeric.size < 2:
62
+ return None
63
+
64
+ sample = _sample_values(numeric)
65
+ vmin = float(np.min(sample))
66
+ vmax = float(np.max(sample))
67
+ if vmin == vmax:
68
+ eps = 1e-3 if abs(vmin) < 1 else abs(vmin) * 1e-3
69
+ support = np.array([vmin - eps, vmin, vmin + eps], dtype=float)
70
+ density = np.array([0.0, 1.0, 0.0], dtype=float)
71
+ return support, density
72
+
73
+ support = np.linspace(vmin, vmax, 100)
74
+ try:
75
+ kernel = stats.gaussian_kde(sample)
76
+ density = kernel(support)
77
+ return support, density
78
+ except Exception:
79
+ bins = int(np.sqrt(sample.size))
80
+ bins = max(5, min(20, bins))
81
+ counts, bin_edges = np.histogram(sample, bins=bins, density=True)
82
+ centers = (bin_edges[:-1] + bin_edges[1:]) / 2.0
83
+ if centers.size == 0:
84
+ return None
85
+ return centers, counts
86
+
87
+
88
+ def _extract_density_data(value: Any) -> tuple[np.ndarray, np.ndarray] | None:
89
+ if _is_missing(value):
90
+ return None
91
+ if isinstance(value, dict):
92
+ if 'density' in value and 'support' in value:
93
+ try:
94
+ density = np.asarray(value['density'], dtype=float).ravel()
95
+ support = np.asarray(value['support'], dtype=float).ravel()
96
+ except (TypeError, ValueError):
97
+ return None
98
+ if density.size == 0 or support.size == 0:
99
+ return None
100
+ if density.size != support.size:
101
+ return None
102
+ mask = np.isfinite(density) & np.isfinite(support)
103
+ if not np.any(mask):
104
+ return None
105
+ density = density[mask]
106
+ support = support[mask]
107
+ order = np.argsort(support)
108
+ return support[order], density[order]
109
+ if 'values' in value:
110
+ return _kde_from_values(np.asarray(value['values']))
111
+ if isinstance(value, (list, tuple, np.ndarray, pd.Series)):
112
+ if isinstance(value, (list, tuple)) and len(value) == 2:
113
+ first = np.asarray(value[0])
114
+ second = np.asarray(value[1])
115
+ if first.ndim == 1 and second.ndim == 1 and first.size == second.size:
116
+ return first.astype(float), second.astype(float)
117
+ return _kde_from_values(np.asarray(value))
118
+ return None
119
+
120
+
121
+ class DensityPlot(StatisticPlot):
122
+ """[PLOT] Density Plot."""
123
+
124
+ name: str = "Density Plot"
125
+ _description: str = textwrap.dedent("""\
126
+ Density plots show smoothed distributions for numeric columns.
127
+ """)
128
+ _description_long: str = textwrap.dedent("""\
129
+ This plot renders kernel density estimates for numeric columns using
130
+ precomputed density statistics when available.
131
+ """)
132
+ refs: list[dict] = []
133
+
134
+ title: str = "Density plot"
135
+ description: str = textwrap.dedent("""\
136
+ The density plot displays kernel density estimates for numeric columns.
137
+ """)
138
+ group_by_feature: bool = True
139
+
140
+ def __str__(self) -> str:
141
+ return 'density'
142
+
143
+ @capture
144
+ def compute(
145
+ self,
146
+ dataframe: pd.DataFrame,
147
+ base_name: str | None = None,
148
+ dataset: Dataset | None = None,
149
+ **kwargs,
150
+ ) -> 'DensityPlot':
151
+ """Compute density plot statistics."""
152
+ self._binary_image = io.BytesIO()
153
+
154
+ if dataframe.empty:
155
+ _plot_placeholder("No statistics available")
156
+ plt.savefig(self._binary_image, format='png')
157
+ return self
158
+
159
+ density_key = str(self)
160
+ if density_key not in dataframe.index:
161
+ _plot_placeholder("Density statistics not available")
162
+ plt.savefig(self._binary_image, format='png')
163
+ return self
164
+
165
+ if dataset is not None:
166
+ numeric_columns = set(dataset.get_columns_names_by_type(DataType.NUMERIC))
167
+ columns_to_show = [col for col in dataframe.columns if col in numeric_columns]
168
+ else:
169
+ columns_to_show = list(dataframe.columns)
170
+
171
+ density_row = dataframe.loc[density_key]
172
+ entries: list[tuple[str, tuple[np.ndarray, np.ndarray]]] = []
173
+ for col in columns_to_show:
174
+ density_data = _extract_density_data(density_row.get(col))
175
+ if density_data is not None:
176
+ entries.append((col, density_data))
177
+
178
+ if not entries:
179
+ _plot_placeholder("No numeric density statistics available")
180
+ plt.savefig(self._binary_image, format='png')
181
+ return self
182
+
183
+ n_plots = len(entries)
184
+ n_cols = 1 if n_plots == 1 else 2
185
+ n_rows = int(np.ceil(n_plots / n_cols))
186
+ fig, axes = plt.subplots(n_rows, n_cols, figsize=(6 * n_cols, 3.5 * n_rows))
187
+ axes_list = np.atleast_1d(axes).ravel()
188
+
189
+ for ax, (col, (support, density)) in zip(axes_list, entries):
190
+ if support.size == 0 or density.size == 0 or support.size != density.size:
191
+ ax.text(0.5, 0.5, "Invalid density data", ha='center', va='center')
192
+ ax.axis('off')
193
+ continue
194
+ ax.plot(support, density, color='tab:blue')
195
+ ax.fill_between(support, density, alpha=0.3, color='tab:blue')
196
+ ax.set_title(_column_label(col, base_name))
197
+ ax.set_xlabel('Value')
198
+ ax.set_ylabel('Density')
199
+
200
+ for ax in axes_list[len(entries):]:
201
+ ax.axis('off')
202
+
203
+ if base_name:
204
+ fig.suptitle(f"Density plot: {base_name}")
205
+ plt.tight_layout(rect=(0, 0, 1, 0.95))
206
+ else:
207
+ plt.tight_layout()
208
+
209
+ plt.savefig(self._binary_image, format='png')
210
+ return self
@@ -0,0 +1,179 @@
1
+ """[PLOT] Histogram 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: object) -> 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 _column_label(column: str, base_name: str | None) -> str:
26
+ column_str = str(column)
27
+ base_name_str = str(base_name) if base_name is not None else None
28
+ if base_name_str:
29
+ if column_str == base_name_str:
30
+ return 'all'
31
+ prefix = f"{base_name_str}_"
32
+ if column_str.startswith(prefix):
33
+ return column_str[len(prefix):]
34
+ if column_str.endswith('_all'):
35
+ return 'all'
36
+ return column_str
37
+
38
+
39
+ def _plot_placeholder(message: str) -> None:
40
+ plt.figure()
41
+ plt.text(0.5, 0.5, message, ha='center', va='center')
42
+ plt.axis('off')
43
+
44
+
45
+ def _histogram_from_values(values: np.ndarray) -> tuple[np.ndarray, np.ndarray] | None:
46
+ if values.size == 0:
47
+ return None
48
+ try:
49
+ numeric = values.astype(float)
50
+ except (ValueError, TypeError):
51
+ return None
52
+ numeric = numeric[np.isfinite(numeric)]
53
+ if numeric.size == 0:
54
+ return None
55
+ bins = int(np.sqrt(numeric.size))
56
+ bins = max(5, min(20, bins))
57
+ counts, bin_edges = np.histogram(numeric, bins=bins)
58
+ return counts, bin_edges
59
+
60
+
61
+ def _extract_hist_data(value: Any) -> tuple[np.ndarray, np.ndarray] | None:
62
+ if _is_missing(value):
63
+ return None
64
+ if isinstance(value, dict):
65
+ if 'counts' in value and ('bins' in value or 'bin_edges' in value):
66
+ counts = np.asarray(value['counts'])
67
+ bins = np.asarray(value.get('bins', value.get('bin_edges')))
68
+ return counts, bins
69
+ if 'hist' in value and ('bins' in value or 'bin_edges' in value):
70
+ counts = np.asarray(value['hist'])
71
+ bins = np.asarray(value.get('bins', value.get('bin_edges')))
72
+ return counts, bins
73
+ if 'values' in value:
74
+ return _histogram_from_values(np.asarray(value['values']))
75
+ if isinstance(value, (list, tuple, np.ndarray, pd.Series)):
76
+ if isinstance(value, (list, tuple)) and len(value) == 2:
77
+ first = np.asarray(value[0])
78
+ second = np.asarray(value[1])
79
+ if first.ndim == 1 and second.ndim == 1:
80
+ if first.size == second.size + 1:
81
+ return second, first
82
+ if second.size == first.size + 1:
83
+ return first, second
84
+ if first.size == second.size and first.size > 0:
85
+ bins = np.arange(first.size + 1)
86
+ return first, bins
87
+ return _histogram_from_values(np.asarray(value))
88
+ return None
89
+
90
+
91
+ class HistogramPlot(StatisticPlot):
92
+ """[PLOT] Histogram Plot."""
93
+
94
+ name: str = "Histogram Plot"
95
+ _description: str = textwrap.dedent("""\
96
+ Histograms show distributions for numeric columns.
97
+ """)
98
+ _description_long: str = textwrap.dedent("""\
99
+ This plot renders a histogram per numeric column to visualize the distribution
100
+ of values. It relies on precomputed histogram statistics when available.
101
+ """)
102
+ refs: list[dict] = []
103
+
104
+ title: str = "Histogram plot"
105
+ description: str = textwrap.dedent("""\
106
+ The histogram plot displays distributions for numeric columns.
107
+ """)
108
+ group_by_feature: bool = True
109
+
110
+ def __str__(self) -> str:
111
+ return 'histogram'
112
+
113
+ @capture
114
+ def compute(
115
+ self,
116
+ dataframe: pd.DataFrame,
117
+ base_name: str | None = None,
118
+ dataset: Dataset | None = None,
119
+ **kwargs) -> 'HistogramPlot':
120
+ """Compute histogram plot statistics."""
121
+ self._binary_image = io.BytesIO()
122
+
123
+ if dataframe.empty:
124
+ _plot_placeholder("No statistics available")
125
+ plt.savefig(self._binary_image, format='png')
126
+ return self
127
+
128
+ hist_key = str(self)
129
+ if hist_key not in dataframe.index:
130
+ _plot_placeholder("Histogram statistics not available")
131
+ plt.savefig(self._binary_image, format='png')
132
+ return self
133
+
134
+ if dataset is not None:
135
+ numeric_columns = set(dataset.get_columns_names_by_type(DataType.NUMERIC))
136
+ columns_to_show = [col for col in dataframe.columns if col in numeric_columns]
137
+ else:
138
+ columns_to_show = list(dataframe.columns)
139
+
140
+ histogram_row = dataframe.loc[hist_key]
141
+ entries: list[tuple[str, tuple[np.ndarray, np.ndarray]]] = []
142
+ for col in columns_to_show:
143
+ hist_data = _extract_hist_data(histogram_row.get(col))
144
+ if hist_data is not None:
145
+ entries.append((col, hist_data))
146
+
147
+ if not entries:
148
+ _plot_placeholder("No numeric histogram statistics available")
149
+ plt.savefig(self._binary_image, format='png')
150
+ return self
151
+
152
+ n_plots = len(entries)
153
+ n_cols = 1 if n_plots == 1 else 2
154
+ n_rows = int(np.ceil(n_plots / n_cols))
155
+ fig, axes = plt.subplots(n_rows, n_cols, figsize=(6 * n_cols, 3.5 * n_rows))
156
+ axes_list = np.atleast_1d(axes).ravel()
157
+
158
+ for ax, (col, (counts, bins)) in zip(axes_list, entries):
159
+ if bins.size != counts.size + 1:
160
+ ax.text(0.5, 0.5, "Invalid histogram data", ha='center', va='center')
161
+ ax.axis('off')
162
+ continue
163
+ widths = np.diff(bins)
164
+ ax.bar(bins[:-1], counts, width=widths, align='edge', edgecolor='black')
165
+ ax.set_title(_column_label(col, base_name))
166
+ ax.set_xlabel('Value')
167
+ ax.set_ylabel('Count')
168
+
169
+ for ax in axes_list[len(entries):]:
170
+ ax.axis('off')
171
+
172
+ if base_name:
173
+ fig.suptitle(f"Histogram: {base_name}")
174
+ plt.tight_layout(rect=(0, 0, 1, 0.95))
175
+ else:
176
+ plt.tight_layout()
177
+
178
+ plt.savefig(self._binary_image, format='png')
179
+ return self