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