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