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,89 @@
1
+ """
2
+ [PLOT] Kaplan-Meier Model Comparison Survival Plot using sksurv
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 numpy as np
11
+ import matplotlib.pyplot as plt
12
+ from sksurv.nonparametric import kaplan_meier_estimator
13
+ from ..metric_plot import MetricPlot, capture
14
+ from ..dataset import Dataset
15
+ if TYPE_CHECKING:
16
+ from ..iaml_pipeline import IAMLPipeline
17
+
18
+
19
+ class KaplanMeierModelComparisonPlot(MetricPlot):
20
+ """[PLOT] Kaplan-Meier Model Comparison Survival"""
21
+
22
+ title: str = "Kaplan-Meier Model Comparison"
23
+ description: str = textwrap.dedent("""
24
+ The Kaplan-Meier Model Comparison Plot is a diagnostic tool used to evaluate the performance of
25
+ survival models by comparing predicted survival curves against the observed survival data.
26
+
27
+ This plot is particularly useful for assessing how well a model can predict the time-to-event
28
+ outcome, such as time until death, disease recurrence, or failure. The observed Kaplan-Meier
29
+ survival curve represents the true survival probability over time, while the model's predicted
30
+ survival curves show the model's estimations.
31
+
32
+ The x-axis represents time, and the y-axis represents the survival probability. Ideally,
33
+ the model-predicted survival curves should closely align with the observed Kaplan-Meier
34
+ curve, indicating good model performance. Discrepancies between the two curves highlight
35
+ areas where the model's predictions diverge from reality, signaling potential issues with
36
+ the model's predictive ability.
37
+
38
+ Additionally, if possible, a Cox proportional hazards model is also trained to serve as
39
+ a baseline. This allows for a better understanding of model performances, as the Cox model
40
+ is a widely-used, interpretable model in survival analysis. By comparing more complex models
41
+ to this baseline, it becomes easier to gauge the improvement (or lack thereof) in predictive
42
+ accuracy.
43
+ """)
44
+
45
+ @capture
46
+ def compute( # pylint: disable=too-many-positional-arguments
47
+ self,
48
+ estimator: IAMLPipeline,
49
+ X: pd.DataFrame,
50
+ y: pd.Series,
51
+ X_train: pd.DataFrame = None,
52
+ y_train: pd.Series=None,
53
+ **kwargs) -> MetricPlot:
54
+ self._binary_image = io.BytesIO()
55
+
56
+ X_train, y_train = Dataset.fix_survival(X_train, y_train)
57
+ X, y = Dataset.fix_survival(X, y)
58
+
59
+ # Fit the Kaplan-Meier model on observed data
60
+ event, time = zip(*y)
61
+
62
+ # Observed data
63
+ time, survival_prob = kaplan_meier_estimator(event, time)
64
+ plt.step(time, survival_prob, where="post", label="Observed", color='blue')
65
+
66
+ # Current model
67
+ survival_predictions = estimator.predict_survival_function(X)
68
+
69
+ mean_survival_prob = np.mean([fn.y for fn in survival_predictions], axis=0)
70
+ mean_survival_time = survival_predictions[0].x
71
+
72
+ plt.step(mean_survival_time, mean_survival_prob,
73
+ where="post", label="Model prediction", color="green")
74
+
75
+ # Customize and save the plot
76
+ plt.title("Kaplan-Meier Curve vs Model Predicted Survival")
77
+ plt.xlabel("Time")
78
+ plt.ylabel("Survival Probability")
79
+ plt.ylim([0, 1])
80
+ plt.legend()
81
+
82
+ plt.savefig(self._binary_image, format='png')
83
+ plt.close()
84
+
85
+ return self
86
+
87
+ @classmethod
88
+ def suitable(cls, type_of_target: str) -> bool:
89
+ return type_of_target == 'survival'
@@ -0,0 +1,70 @@
1
+ """[PLOT] Line plot for descriptive statistics."""
2
+ from __future__ import annotations
3
+
4
+ import io
5
+ import textwrap
6
+ import pandas as pd
7
+ import matplotlib.pyplot as plt
8
+
9
+ from ..plot import StatisticPlot, capture
10
+
11
+
12
+ def _plot_placeholder(message: str) -> None:
13
+ plt.figure()
14
+ plt.text(0.5, 0.5, message, ha='center', va='center')
15
+ plt.axis('off')
16
+
17
+
18
+ class LinePlot(StatisticPlot):
19
+ """[PLOT] Line Plot."""
20
+
21
+ title: str = "Line plot"
22
+ description: str = textwrap.dedent("""\
23
+ The line plot compares the null count and count statistics across columns.
24
+ """)
25
+ group_by_feature: bool = False
26
+
27
+ @capture
28
+ def compute(self, dataframe: pd.DataFrame, **kwargs) -> 'LinePlot':
29
+ """Compute line plot statistics."""
30
+ self._binary_image = io.BytesIO()
31
+
32
+ if dataframe.empty:
33
+ _plot_placeholder("No statistics available")
34
+ plt.savefig(self._binary_image, format='png')
35
+ return self
36
+
37
+ if 'count' not in dataframe.index or 'null_count' not in dataframe.index:
38
+ _plot_placeholder("Count statistics not available")
39
+ plt.savefig(self._binary_image, format='png')
40
+ return self
41
+
42
+ columns_to_show = [
43
+ col for col in dataframe.columns
44
+ if isinstance(col, str) and col.endswith('_all')
45
+ ]
46
+ if columns_to_show:
47
+ labels = [col[:-4] for col in columns_to_show]
48
+ else:
49
+ columns_to_show = list(dataframe.columns)
50
+ labels = [str(col) for col in columns_to_show]
51
+
52
+ counts = dataframe.loc['count', columns_to_show]
53
+ null_counts = dataframe.loc['null_count', columns_to_show]
54
+ if counts.empty or null_counts.empty:
55
+ _plot_placeholder("Count statistics not available")
56
+ plt.savefig(self._binary_image, format='png')
57
+ return self
58
+
59
+ plt.figure()
60
+ plt.plot(labels, null_counts.values, marker='o', label='Null Count', color='red')
61
+ plt.plot(labels, counts.values, marker='s', label='Count', color='blue')
62
+ plt.title('Null Count vs Count')
63
+ plt.xlabel('Columns')
64
+ plt.ylabel('Value')
65
+ plt.legend()
66
+ plt.xticks(rotation=45)
67
+ plt.tight_layout()
68
+
69
+ plt.savefig(self._binary_image, format='png')
70
+ return self
@@ -0,0 +1,203 @@
1
+ """[PLOT] Missingness 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 _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 _missing_row_from_stats(dataframe: pd.DataFrame) -> pd.Series | None:
24
+ if dataframe.empty:
25
+ return None
26
+ if 'missing_rate' in dataframe.index:
27
+ row = dataframe.loc['missing_rate']
28
+ elif dataframe.shape[0] == 1:
29
+ row = dataframe.iloc[0]
30
+ else:
31
+ return None
32
+ if isinstance(row, pd.DataFrame):
33
+ if row.empty:
34
+ return None
35
+ row = row.iloc[0]
36
+ return row
37
+
38
+
39
+ def _split_missing_column(column: Any) -> tuple[str, str]:
40
+ column_str = str(column)
41
+ if column_str.endswith('_all'):
42
+ return column_str[:-4], 'all'
43
+ if '_' in column_str:
44
+ base, label = column_str.rsplit('_', 1)
45
+ return base, label
46
+ return column_str, 'all'
47
+
48
+
49
+ def _missing_frame_from_stats(
50
+ dataframe: pd.DataFrame,
51
+ dataset: Dataset | None,
52
+ ) -> pd.DataFrame | None:
53
+ row = _missing_row_from_stats(dataframe)
54
+ if row is None:
55
+ return None
56
+
57
+ if dataset is not None and dataset.type_of_target != 'survival':
58
+ columns = [
59
+ col for col in dataset.X.columns
60
+ if dataset.columns_types[col][1] not in (DataType.TEXT, DataType.SHORT_TEXT)
61
+ ]
62
+ if not columns:
63
+ return None
64
+ if dataset.type_of_target == 'continuous':
65
+ values = []
66
+ for col in columns:
67
+ values.append(row.get(col, np.nan))
68
+ return pd.DataFrame([values], index=['all'], columns=columns)
69
+
70
+ class_labels = list(pd.unique(dataset.y))
71
+ labels = ['all'] + [str(label) for label in class_labels]
72
+ data = np.full((len(labels), len(columns)), np.nan, dtype=float)
73
+ for col_idx, col in enumerate(columns):
74
+ for row_idx, label in enumerate(labels):
75
+ key = f"{col}_{label}"
76
+ data[row_idx, col_idx] = row.get(key, np.nan)
77
+ return pd.DataFrame(data, index=labels, columns=columns)
78
+
79
+ columns: list[str] = []
80
+ labels: list[str] = []
81
+ values: dict[tuple[str, str], float] = {}
82
+ for col_name, value in row.items():
83
+ base, label = _split_missing_column(col_name)
84
+ if base not in columns:
85
+ columns.append(base)
86
+ if label not in labels:
87
+ labels.append(label)
88
+ values[(label, base)] = value
89
+
90
+ if not columns or not labels:
91
+ return None
92
+
93
+ if 'all' in labels:
94
+ labels = ['all'] + [label for label in labels if label != 'all']
95
+
96
+ data = np.full((len(labels), len(columns)), np.nan, dtype=float)
97
+ for row_idx, label in enumerate(labels):
98
+ for col_idx, base in enumerate(columns):
99
+ data[row_idx, col_idx] = values.get((label, base), np.nan)
100
+ return pd.DataFrame(data, index=labels, columns=columns)
101
+
102
+
103
+ def _missing_frame_from_dataset(dataset: Dataset) -> pd.DataFrame | None:
104
+ if dataset.type_of_target == 'survival':
105
+ return None
106
+
107
+ columns = [
108
+ col for col in dataset.X.columns
109
+ if dataset.columns_types[col][1] not in (DataType.TEXT, DataType.SHORT_TEXT)
110
+ ]
111
+ if not columns:
112
+ return None
113
+
114
+ frame = dataset.X[columns]
115
+ overall = frame.isna().mean(axis=0)
116
+
117
+ if dataset.type_of_target == 'continuous':
118
+ return pd.DataFrame([overall.to_numpy()], index=['all'], columns=columns)
119
+
120
+ grouped = frame.isna().groupby(dataset.y, sort=False).mean()
121
+ labels = list(pd.unique(dataset.y))
122
+ grouped = grouped.reindex(labels)
123
+ grouped.index = [str(label) for label in grouped.index]
124
+
125
+ overall_frame = pd.DataFrame([overall.to_numpy()], index=['all'], columns=columns)
126
+ return pd.concat([overall_frame, grouped])
127
+
128
+
129
+ def _figure_size(n_cols: int, n_rows: int) -> tuple[float, float]:
130
+ width = float(min(14.0, max(5.0, 0.6 * n_cols + 2.5)))
131
+ height = float(min(10.0, max(3.5, 0.5 * n_rows + 2.0)))
132
+ return width, height
133
+
134
+
135
+ class MissingnessHeatmapPlot(StatisticPlot):
136
+ """[PLOT] Missingness Heatmap Plot."""
137
+
138
+ name: str = "Missingness Heatmap"
139
+ _description: str = textwrap.dedent("""\
140
+ Missingness heatmaps summarize the percentage of missing values.
141
+ """)
142
+ _description_long: str = textwrap.dedent("""\
143
+ This plot displays a heatmap of missing rates per feature, optionally
144
+ broken down by class labels for classification datasets.
145
+ """)
146
+ refs: list[dict] = []
147
+
148
+ title: str = "Missingness heatmap"
149
+ description: str = textwrap.dedent("""\
150
+ The missingness heatmap shows missing value rates per feature.
151
+ """)
152
+ group_by_feature: bool = False
153
+
154
+ def __str__(self) -> str:
155
+ return 'missing_rate'
156
+
157
+ @capture
158
+ def compute(
159
+ self,
160
+ dataframe: pd.DataFrame,
161
+ dataset: Dataset | None = None,
162
+ base_name: str | None = None,
163
+ **kwargs,
164
+ ) -> 'MissingnessHeatmapPlot':
165
+ """Compute missingness heatmap statistics."""
166
+ self._binary_image = io.BytesIO()
167
+
168
+ missing_df = _missing_frame_from_stats(dataframe, dataset)
169
+ if missing_df is None and dataset is not None:
170
+ missing_df = _missing_frame_from_dataset(dataset)
171
+
172
+ if missing_df is None or missing_df.empty:
173
+ _plot_placeholder("Missingness statistics not available")
174
+ plt.savefig(self._binary_image, format='png')
175
+ return self
176
+
177
+ values = missing_df.to_numpy(dtype=float)
178
+ masked = np.ma.masked_invalid(values)
179
+ n_rows, n_cols = masked.shape
180
+ fig_size = _figure_size(n_cols, n_rows)
181
+ fig, ax = plt.subplots(figsize=fig_size)
182
+
183
+ image = ax.imshow(masked, cmap='Reds', vmin=0, vmax=1, aspect='auto')
184
+ plt.colorbar(image, ax=ax, fraction=0.046, pad=0.04, label='Missing rate')
185
+
186
+ ax.set_xticks(np.arange(n_cols))
187
+ ax.set_yticks(np.arange(n_rows))
188
+ ax.set_xticklabels([str(label) for label in missing_df.columns], rotation=45, ha='right')
189
+ ax.set_yticklabels([str(label) for label in missing_df.index])
190
+
191
+ title = "Missingness heatmap"
192
+ if base_name:
193
+ title = f"Missingness heatmap: {base_name}"
194
+ ax.set_title(title)
195
+ ax.set_xlabel('Features')
196
+ ax.set_ylabel('Group')
197
+
198
+ tick_size = 10 if n_cols <= 12 else 8 if n_cols <= 20 else 6
199
+ ax.tick_params(axis='both', which='major', labelsize=tick_size)
200
+
201
+ plt.tight_layout()
202
+ plt.savefig(self._binary_image, format='png')
203
+ return self
@@ -0,0 +1,217 @@
1
+ """[PLOT] Outlier 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 _outlier_row_from_stats(dataframe: pd.DataFrame, key: str) -> pd.Series | None:
24
+ if dataframe.empty:
25
+ return None
26
+ if key in dataframe.index:
27
+ row = dataframe.loc[key]
28
+ elif dataframe.shape[0] == 1:
29
+ row = dataframe.iloc[0]
30
+ else:
31
+ return None
32
+ if isinstance(row, pd.DataFrame):
33
+ if row.empty:
34
+ return None
35
+ row = row.iloc[0]
36
+ return row
37
+
38
+
39
+ def _split_outlier_column(column: Any) -> tuple[str, str]:
40
+ column_str = str(column)
41
+ if column_str.endswith('_all'):
42
+ return column_str[:-4], 'all'
43
+ if '_' in column_str:
44
+ base, label = column_str.rsplit('_', 1)
45
+ return base, label
46
+ return column_str, 'all'
47
+
48
+
49
+ def _outlier_frame_from_stats(
50
+ dataframe: pd.DataFrame,
51
+ dataset: Dataset | None,
52
+ key: str,
53
+ ) -> pd.DataFrame | None:
54
+ row = _outlier_row_from_stats(dataframe, key)
55
+ if row is None:
56
+ return None
57
+
58
+ if dataset is not None and dataset.type_of_target != 'survival':
59
+ numeric_columns = list(dataset.get_columns_names_by_type(DataType.NUMERIC))
60
+ if not numeric_columns:
61
+ return None
62
+ if dataset.type_of_target == 'continuous':
63
+ values = [row.get(col, np.nan) for col in numeric_columns]
64
+ return pd.DataFrame([values], index=['all'], columns=numeric_columns)
65
+
66
+ class_labels = list(pd.unique(dataset.y))
67
+ labels = ['all'] + [str(label) for label in class_labels]
68
+ data = np.full((len(labels), len(numeric_columns)), np.nan, dtype=float)
69
+ for col_idx, col in enumerate(numeric_columns):
70
+ for row_idx, label in enumerate(labels):
71
+ key_name = f"{col}_{label}"
72
+ data[row_idx, col_idx] = row.get(key_name, np.nan)
73
+ return pd.DataFrame(data, index=labels, columns=numeric_columns)
74
+
75
+ columns: list[str] = []
76
+ labels: list[str] = []
77
+ values: dict[tuple[str, str], float] = {}
78
+ for col_name, value in row.items():
79
+ base, label = _split_outlier_column(col_name)
80
+ if base not in columns:
81
+ columns.append(base)
82
+ if label not in labels:
83
+ labels.append(label)
84
+ values[(label, base)] = value
85
+
86
+ if not columns or not labels:
87
+ return None
88
+
89
+ if 'all' in labels:
90
+ labels = ['all'] + [label for label in labels if label != 'all']
91
+
92
+ data = np.full((len(labels), len(columns)), np.nan, dtype=float)
93
+ for row_idx, label in enumerate(labels):
94
+ for col_idx, base in enumerate(columns):
95
+ data[row_idx, col_idx] = values.get((label, base), np.nan)
96
+ return pd.DataFrame(data, index=labels, columns=columns)
97
+
98
+
99
+ def _figure_size(n_cols: int) -> tuple[float, float]:
100
+ width = float(min(14.0, max(6.0, 0.7 * n_cols + 2.5)))
101
+ height = 4.5 if n_cols <= 8 else 5.5
102
+ return width, height
103
+
104
+
105
+ class OutlierPlot(StatisticPlot):
106
+ """[PLOT] Outlier Plot."""
107
+
108
+ name: str = "Outlier Plot"
109
+ _description: str = textwrap.dedent("""\
110
+ Outlier plots show counts of outliers per column.
111
+ """)
112
+ _description_long: str = textwrap.dedent("""\
113
+ This plot visualizes outlier counts based on the 1.5*IQR rule for each
114
+ numeric column, optionally broken down by class labels.
115
+ """)
116
+ refs: list[dict] = []
117
+
118
+ title: str = "Outlier plot"
119
+ description: str = textwrap.dedent("""\
120
+ The outlier plot shows outlier counts per column using box/strip visuals.
121
+ """)
122
+ group_by_feature: bool = False
123
+
124
+ def __str__(self) -> str:
125
+ return 'outlier_count_iqr'
126
+
127
+ @capture
128
+ def compute(
129
+ self,
130
+ dataframe: pd.DataFrame,
131
+ dataset: Dataset | None = None,
132
+ base_name: str | None = None,
133
+ **kwargs,
134
+ ) -> 'OutlierPlot':
135
+ """Compute outlier plot statistics."""
136
+ self._binary_image = io.BytesIO()
137
+
138
+ outlier_key = str(self)
139
+ outlier_df = _outlier_frame_from_stats(dataframe, dataset, outlier_key)
140
+ if outlier_df is None or outlier_df.empty:
141
+ _plot_placeholder("Outlier statistics not available")
142
+ plt.savefig(self._binary_image, format='png')
143
+ return self
144
+
145
+ columns = [col for col in outlier_df.columns]
146
+ if not columns:
147
+ _plot_placeholder("Outlier statistics not available")
148
+ plt.savefig(self._binary_image, format='png')
149
+ return self
150
+
151
+ values_per_column: list[np.ndarray] = []
152
+ columns_to_plot: list[str] = []
153
+ for col in columns:
154
+ values = pd.to_numeric(outlier_df[col], errors='coerce').to_numpy(dtype=float)
155
+ values = values[np.isfinite(values)]
156
+ if values.size == 0:
157
+ continue
158
+ columns_to_plot.append(col)
159
+ values_per_column.append(values)
160
+
161
+ if not columns_to_plot:
162
+ _plot_placeholder("No numeric outlier statistics available")
163
+ plt.savefig(self._binary_image, format='png')
164
+ return self
165
+
166
+ fig, ax = plt.subplots(figsize=_figure_size(len(columns_to_plot)))
167
+ positions = np.arange(1, len(columns_to_plot) + 1)
168
+ ax.boxplot(
169
+ values_per_column,
170
+ positions=positions,
171
+ widths=0.55,
172
+ showfliers=False,
173
+ patch_artist=True,
174
+ boxprops={'facecolor': '#d9d9d9', 'edgecolor': '#555555'},
175
+ medianprops={'color': '#333333'},
176
+ )
177
+
178
+ labels = list(outlier_df.index)
179
+ n_labels = len(labels)
180
+ offsets = np.linspace(-0.18, 0.18, n_labels) if n_labels > 1 else np.array([0.0])
181
+ cmap = plt.get_cmap('tab10')
182
+ colors = cmap(np.linspace(0, 1, max(1, n_labels)))
183
+
184
+ for label_idx, label in enumerate(labels):
185
+ row_values = pd.to_numeric(
186
+ outlier_df.loc[label, columns_to_plot],
187
+ errors='coerce',
188
+ ).to_numpy(dtype=float)
189
+ mask = np.isfinite(row_values)
190
+ if not np.any(mask):
191
+ continue
192
+ x_positions = positions[mask] + offsets[label_idx]
193
+ y_values = row_values[mask]
194
+ ax.scatter(
195
+ x_positions,
196
+ y_values,
197
+ s=28,
198
+ color=colors[label_idx],
199
+ alpha=0.85,
200
+ label=str(label),
201
+ zorder=3,
202
+ )
203
+
204
+ ax.set_xticks(positions)
205
+ ax.set_xticklabels([str(col) for col in columns_to_plot], rotation=30, ha='right')
206
+ ax.set_ylabel('Outlier count')
207
+ title = "Outlier count (IQR)"
208
+ if base_name:
209
+ title = f"Outlier count (IQR): {base_name}"
210
+ ax.set_title(title)
211
+
212
+ if n_labels > 1:
213
+ ax.legend(title='Group', fontsize=8)
214
+
215
+ plt.tight_layout()
216
+ plt.savefig(self._binary_image, format='png')
217
+ return self