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.
- iaml/__init__.py +56 -0
- iaml/actionable.py +11 -0
- iaml/actionables/__init__.py +21 -0
- iaml/actionables/boosting/__init__.py +4 -0
- iaml/actionables/boosting/act_adaboost.py +59 -0
- iaml/actionables/cleaning/__init__.py +26 -0
- iaml/actionables/cleaning/act_categorical_imputer.py +124 -0
- iaml/actionables/cleaning/act_count_vectorizer.py +204 -0
- iaml/actionables/cleaning/act_drop_categorical_column.py +51 -0
- iaml/actionables/cleaning/act_drop_date_column.py +48 -0
- iaml/actionables/cleaning/act_drop_high_cardinality_categorical.py +337 -0
- iaml/actionables/cleaning/act_drop_numerical_column.py +75 -0
- iaml/actionables/cleaning/act_drop_textual_column.py +51 -0
- iaml/actionables/cleaning/act_encode_target_column.py +56 -0
- iaml/actionables/cleaning/act_frequency_encoder.py +127 -0
- iaml/actionables/cleaning/act_hashing_vectorizer.py +186 -0
- iaml/actionables/cleaning/act_knn_imputer.py +152 -0
- iaml/actionables/cleaning/act_mean_column.py +79 -0
- iaml/actionables/cleaning/act_mice.py +464 -0
- iaml/actionables/cleaning/act_missing_count_feature.py +109 -0
- iaml/actionables/cleaning/act_missing_indicator.py +124 -0
- iaml/actionables/cleaning/act_onehot.py +65 -0
- iaml/actionables/cleaning/act_ordinal_encoder.py +177 -0
- iaml/actionables/cleaning/act_rare_category_grouper.py +173 -0
- iaml/actionables/cleaning/act_simple_imputer.py +109 -0
- iaml/actionables/cleaning/act_split_date.py +68 -0
- iaml/actionables/cleaning/act_target_encoder.py +274 -0
- iaml/actionables/cleaning/act_text_normalizer.py +241 -0
- iaml/actionables/cleaning/act_tf_idf.py +80 -0
- iaml/actionables/cleaning/act_word2vec.py +150 -0
- iaml/actionables/features_precleaning/__init__.py +12 -0
- iaml/actionables/features_precleaning/act_coerce_numeric_strings.py +194 -0
- iaml/actionables/features_precleaning/act_date_converter.py +99 -0
- iaml/actionables/features_precleaning/act_drop_bad_quality_rows.py +77 -0
- iaml/actionables/features_precleaning/act_drop_duplicate_rows.py +131 -0
- iaml/actionables/features_precleaning/act_drop_high_missing_columns.py +94 -0
- iaml/actionables/features_precleaning/act_drop_id_like_columns.py +294 -0
- iaml/actionables/features_precleaning/act_normalize_column_names.py +157 -0
- iaml/actionables/features_precleaning/act_sentinel_to_na_n.py +270 -0
- iaml/actionables/features_precleaning/act_trim_space.py +79 -0
- iaml/actionables/features_preprocessing/__init__.py +18 -0
- iaml/actionables/features_preprocessing/act_cyclical_date_encoding.py +212 -0
- iaml/actionables/features_preprocessing/act_fast_ica.py +161 -0
- iaml/actionables/features_preprocessing/act_feature_agglomeration.py +90 -0
- iaml/actionables/features_preprocessing/act_k_bins_discretizer.py +207 -0
- iaml/actionables/features_preprocessing/act_k_means_features.py +296 -0
- iaml/actionables/features_preprocessing/act_kernel_pca.py +143 -0
- iaml/actionables/features_preprocessing/act_log_transformer.py +122 -0
- iaml/actionables/features_preprocessing/act_nystroem.py +100 -0
- iaml/actionables/features_preprocessing/act_pca.py +77 -0
- iaml/actionables/features_preprocessing/act_polynomial_features.py +86 -0
- iaml/actionables/features_preprocessing/act_power_transformer.py +106 -0
- iaml/actionables/features_preprocessing/act_quantile_transformer.py +114 -0
- iaml/actionables/features_preprocessing/act_rbf_sampler.py +88 -0
- iaml/actionables/features_preprocessing/act_select_percentile.py +112 -0
- iaml/actionables/features_preprocessing/act_sparse_random_projection.py +157 -0
- iaml/actionables/features_preprocessing/act_truncated_svd.py +137 -0
- iaml/actionables/features_selection/__init__.py +8 -0
- iaml/actionables/features_selection/act_permutation_importance_selector.py +421 -0
- iaml/actionables/features_selection/act_remove_high_correlated_column.py +70 -0
- iaml/actionables/features_selection/act_remove_low_variance_column.py +74 -0
- iaml/actionables/features_selection/act_rfe.py +214 -0
- iaml/actionables/features_selection/act_select_from_model.py +325 -0
- iaml/actionables/features_selection/act_select_k_best.py +181 -0
- iaml/actionables/features_selection/act_vif_selector.py +130 -0
- iaml/actionables/imbalance/__init__.py +10 -0
- iaml/actionables/imbalance/act_adasyn.py +150 -0
- iaml/actionables/imbalance/act_borderline_smote.py +171 -0
- iaml/actionables/imbalance/act_near_miss.py +158 -0
- iaml/actionables/imbalance/act_random_over_sampling.py +60 -0
- iaml/actionables/imbalance/act_random_under_sampler.py +135 -0
- iaml/actionables/imbalance/act_smote.py +162 -0
- iaml/actionables/imbalance/act_smote_tomek.py +182 -0
- iaml/actionables/imbalance/act_smoteenn.py +193 -0
- iaml/actionables/imbalance/act_tomek_links.py +138 -0
- iaml/actionables/normalize/__init__.py +6 -0
- iaml/actionables/normalize/act_max_abs_scaler.py +78 -0
- iaml/actionables/normalize/act_minmax_scaler.py +56 -0
- iaml/actionables/normalize/act_normalizer.py +95 -0
- iaml/actionables/normalize/act_robust_scaler.py +111 -0
- iaml/actionables/normalize/act_standard_scaler.py +55 -0
- iaml/actionables/predictors/__init__.py +6 -0
- iaml/actionables/predictors/_xgboost.py +16 -0
- iaml/actionables/predictors/classifier/__init__.py +26 -0
- iaml/actionables/predictors/classifier/act_bagging_classifier.py +113 -0
- iaml/actionables/predictors/classifier/act_bernoulli_nb.py +89 -0
- iaml/actionables/predictors/classifier/act_catboost_classifier.py +135 -0
- iaml/actionables/predictors/classifier/act_complement_nb.py +106 -0
- iaml/actionables/predictors/classifier/act_decision_tree_classifier.py +117 -0
- iaml/actionables/predictors/classifier/act_extra_trees_classifier.py +115 -0
- iaml/actionables/predictors/classifier/act_gaussian_nb.py +53 -0
- iaml/actionables/predictors/classifier/act_hist_gradient_boosting_classifier.py +144 -0
- iaml/actionables/predictors/classifier/act_knn.py +86 -0
- iaml/actionables/predictors/classifier/act_light_gbm_classifier.py +211 -0
- iaml/actionables/predictors/classifier/act_linear_discriminant_analysis.py +63 -0
- iaml/actionables/predictors/classifier/act_linear_svc.py +134 -0
- iaml/actionables/predictors/classifier/act_logistic_regression.py +92 -0
- iaml/actionables/predictors/classifier/act_mlp_classifier.py +107 -0
- iaml/actionables/predictors/classifier/act_multinomial_nb.py +76 -0
- iaml/actionables/predictors/classifier/act_passive_aggressive_classifier.py +141 -0
- iaml/actionables/predictors/classifier/act_quadratic_discriminant_analysis.py +72 -0
- iaml/actionables/predictors/classifier/act_randomforest.py +113 -0
- iaml/actionables/predictors/classifier/act_ridge_classifier.py +116 -0
- iaml/actionables/predictors/classifier/act_sgd_classifier.py +149 -0
- iaml/actionables/predictors/classifier/act_svm_svc.py +88 -0
- iaml/actionables/predictors/classifier/act_xgboost.py +111 -0
- iaml/actionables/predictors/regressor/__init__.py +27 -0
- iaml/actionables/predictors/regressor/act_ada_boost_regressor.py +75 -0
- iaml/actionables/predictors/regressor/act_ard_regression.py +95 -0
- iaml/actionables/predictors/regressor/act_catboost_regressor.py +134 -0
- iaml/actionables/predictors/regressor/act_decision_tree_regressor.py +111 -0
- iaml/actionables/predictors/regressor/act_elastic_net_regressor.py +109 -0
- iaml/actionables/predictors/regressor/act_extra_trees_regressor.py +113 -0
- iaml/actionables/predictors/regressor/act_gaussian_process_regressor.py +55 -0
- iaml/actionables/predictors/regressor/act_gboost_regressor.py +95 -0
- iaml/actionables/predictors/regressor/act_hist_gradient_boosting_regressor.py +105 -0
- iaml/actionables/predictors/regressor/act_huber_regressor.py +101 -0
- iaml/actionables/predictors/regressor/act_knn_regressor.py +86 -0
- iaml/actionables/predictors/regressor/act_lasso_regressor.py +103 -0
- iaml/actionables/predictors/regressor/act_light_gbm_regressor.py +201 -0
- iaml/actionables/predictors/regressor/act_linear_regression.py +43 -0
- iaml/actionables/predictors/regressor/act_mlp_regressor.py +104 -0
- iaml/actionables/predictors/regressor/act_poisson_regressor.py +111 -0
- iaml/actionables/predictors/regressor/act_quantile_regressor.py +87 -0
- iaml/actionables/predictors/regressor/act_randomforest_regressor.py +116 -0
- iaml/actionables/predictors/regressor/act_ransac_regressor.py +106 -0
- iaml/actionables/predictors/regressor/act_ridge_regressor.py +107 -0
- iaml/actionables/predictors/regressor/act_sgd_regressor.py +106 -0
- iaml/actionables/predictors/regressor/act_svm_svr.py +81 -0
- iaml/actionables/predictors/regressor/act_xgboost_regressor.py +97 -0
- iaml/actionables/predictors/survival/__init__.py +12 -0
- iaml/actionables/predictors/survival/act_aalen_additive_model.py +83 -0
- iaml/actionables/predictors/survival/act_cox.py +110 -0
- iaml/actionables/predictors/survival/act_coxnet_survival_analysis.py +134 -0
- iaml/actionables/predictors/survival/act_extra_survival_trees.py +101 -0
- iaml/actionables/predictors/survival/act_fast_survival_svm.py +102 -0
- iaml/actionables/predictors/survival/act_gradient_boosting_survival_analysis.py +93 -0
- iaml/actionables/predictors/survival/act_random_survival_forest.py +91 -0
- iaml/actionables/predictors/survival/act_survival_component_wise_gboost.py +80 -0
- iaml/actionables/predictors/survival/act_survival_tree.py +120 -0
- iaml/actionables/predictors/survival/act_survival_xgboost.py +9 -0
- iaml/actionables/predictors/survival/act_weibull_aft.py +230 -0
- iaml/cache.py +61 -0
- iaml/cache_keys.py +57 -0
- iaml/candidate.py +736 -0
- iaml/core_dispatcher.py +125 -0
- iaml/data_type.py +11 -0
- iaml/dataset.py +506 -0
- iaml/decorators/__init__.py +3 -0
- iaml/decorators/all.py +4 -0
- iaml/decorators/is_step.py +45 -0
- iaml/decorators/runner.py +100 -0
- iaml/explanation.py +112 -0
- iaml/iaml.py +1072 -0
- iaml/iaml_pipeline.py +600 -0
- iaml/logger.py +138 -0
- iaml/meta_explorer_step.py +62 -0
- iaml/meta_ordered_step.py +28 -0
- iaml/meta_partial_explorer_step.py +34 -0
- iaml/meta_singleton.py +24 -0
- iaml/metastep.py +211 -0
- iaml/metric.py +111 -0
- iaml/metric_plot.py +82 -0
- iaml/metrics/__init__.py +21 -0
- iaml/metrics/_classification.py +28 -0
- iaml/metrics/_survival_times.py +22 -0
- iaml/metrics/accuracy_metric.py +59 -0
- iaml/metrics/balanced_accuracy_metric.py +67 -0
- iaml/metrics/brier_score.py +90 -0
- iaml/metrics/classification_error_metric.py +66 -0
- iaml/metrics/concordance_index_ipcw.py +84 -0
- iaml/metrics/concordance_index_metric.py +67 -0
- iaml/metrics/cumulative_dynamic_auc.py +119 -0
- iaml/metrics/f1_score_metric.py +71 -0
- iaml/metrics/integrated_brier_score.py +98 -0
- iaml/metrics/integrated_brier_score_loss.py +41 -0
- iaml/metrics/mean_absolute_error_metric.py +46 -0
- iaml/metrics/mean_squared_error_metric.py +46 -0
- iaml/metrics/mean_squared_log_error_metric.py +49 -0
- iaml/metrics/median_absolute_error_metric.py +48 -0
- iaml/metrics/precision_metric.py +63 -0
- iaml/metrics/r2_score_metric.py +45 -0
- iaml/metrics/recall_metric.py +65 -0
- iaml/metrics/roc_auc_metric.py +50 -0
- iaml/metrics/specificity_metric.py +44 -0
- iaml/metrics/specificity_multiclass_metric.py +55 -0
- iaml/metrics/specificity_multilabel_metric.py +60 -0
- iaml/optimizers/__init__.py +5 -0
- iaml/optimizers/bayesian_optimizer.py +193 -0
- iaml/optimizers/genetic_optimizer.py +284 -0
- iaml/optimizers/optimizer.py +31 -0
- iaml/optimizers/random_optimizer.py +101 -0
- iaml/plot.py +138 -0
- iaml/plots/__init__.py +32 -0
- iaml/plots/bar_plot.py +141 -0
- iaml/plots/box_plot.py +166 -0
- iaml/plots/class_prediction_error_plot.py +37 -0
- iaml/plots/classification_report_plot.py +35 -0
- iaml/plots/confusion_matrix_plot.py +34 -0
- iaml/plots/correlation_heatmap_plot.py +201 -0
- iaml/plots/cumulative_hazard_plot.py +72 -0
- iaml/plots/density_plot.py +210 -0
- iaml/plots/histogram_plot.py +179 -0
- iaml/plots/kaplan_meier_comparison_plot.py +89 -0
- iaml/plots/line_plot.py +70 -0
- iaml/plots/missingness_heatmap_plot.py +203 -0
- iaml/plots/outlier_plot.py +217 -0
- iaml/plots/pair_plot.py +228 -0
- iaml/plots/precision_recall_curve_plot.py +86 -0
- iaml/plots/prediction_error_plot.py +34 -0
- iaml/plots/qq_plot.py +220 -0
- iaml/plots/residual_plot.py +38 -0
- iaml/plots/roc_dynamique_curve_plot.py +79 -0
- iaml/plots/rocauc_plot.py +96 -0
- iaml/plots/shap_plot.py +187 -0
- iaml/plots/target_distribution_plot.py +241 -0
- iaml/plots/violin_plot.py +206 -0
- iaml/predictor.py +139 -0
- iaml/reference.py +65 -0
- iaml/shared_cache.py +90 -0
- iaml/sklearn_preprocessor.py +74 -0
- iaml/splitters/__init__.py +3 -0
- iaml/splitters/kfold_splitter.py +32 -0
- iaml/splitters/random_splitter.py +26 -0
- iaml/stack.py +39 -0
- iaml/statistic.py +66 -0
- iaml/statistics/__init__.py +77 -0
- iaml/statistics/anova_statistic.py +80 -0
- iaml/statistics/cardinality_ratio_statistic.py +63 -0
- iaml/statistics/category_cooccurrence_statistic.py +79 -0
- iaml/statistics/chi_square_statistic.py +81 -0
- iaml/statistics/coef_variation_statistic.py +72 -0
- iaml/statistics/correlation_with_target.py +105 -0
- iaml/statistics/count.py +72 -0
- iaml/statistics/data_type_summary_statistic.py +74 -0
- iaml/statistics/duplicate_row_statistic.py +56 -0
- iaml/statistics/effect_size_statistic.py +129 -0
- iaml/statistics/entropy_statistic.py +69 -0
- iaml/statistics/event_rate_statistic.py +52 -0
- iaml/statistics/grouped_mean_statistic.py +60 -0
- iaml/statistics/iqr_statistic.py +66 -0
- iaml/statistics/kurtosis.py +50 -0
- iaml/statistics/mad_statistic.py +66 -0
- iaml/statistics/mean.py +61 -0
- iaml/statistics/median_statistic.py +61 -0
- iaml/statistics/minmax.py +60 -0
- iaml/statistics/missing_rate_statistic.py +62 -0
- iaml/statistics/mode.py +47 -0
- iaml/statistics/most_frequent_ratio.py +81 -0
- iaml/statistics/outlier_count_iqr_statistic.py +76 -0
- iaml/statistics/quantile.py +59 -0
- iaml/statistics/range.py +53 -0
- iaml/statistics/rare_category_rate.py +92 -0
- iaml/statistics/skewness.py +53 -0
- iaml/statistics/stdev.py +50 -0
- iaml/statistics/summary_table_statistic.py +60 -0
- iaml/statistics/time_by_group_statistic.py +83 -0
- iaml/statistics/time_summary_statistic.py +56 -0
- iaml/statistics/top_k_value_counts.py +68 -0
- iaml/statistics/unique_count_statistic.py +57 -0
- iaml/statistics/value_counts.py +63 -0
- iaml/statistics/variance.py +51 -0
- iaml/statistics/violin.py +63 -0
- iaml/step.py +600 -0
- iaml/step_cache.py +87 -0
- iaml/step_wrapper.py +79 -0
- iaml/timed_pool_executor.py +492 -0
- iaml/type_of_target.py +68 -0
- iaml/void_step.py +101 -0
- iaml/worker_manager.py +169 -0
- iaml/wrapper/__init__.py +4 -0
- iaml/wrapper/wrap_basic_gridsearch.py +68 -0
- iaml/wrapper/wrap_genetic_gridsearch.py +293 -0
- iaml/wrapper/wrap_iterative_gridsearch.py +399 -0
- pyiaml-1.0.0.dist-info/METADATA +802 -0
- pyiaml-1.0.0.dist-info/RECORD +279 -0
- pyiaml-1.0.0.dist-info/WHEEL +5 -0
- pyiaml-1.0.0.dist-info/licenses/LICENSE +674 -0
- 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
|