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,101 @@
|
|
|
1
|
+
|
|
2
|
+
"""
|
|
3
|
+
Base class of IAML Optimizer. Optimizer receive a pool of Candidates,
|
|
4
|
+
optimize parameters and return a new pool of candidate
|
|
5
|
+
"""
|
|
6
|
+
from copy import deepcopy
|
|
7
|
+
import random
|
|
8
|
+
import time
|
|
9
|
+
from ..candidate import Candidate
|
|
10
|
+
from .optimizer import Optimizer
|
|
11
|
+
from ..step import Step
|
|
12
|
+
from ..logger import Logger
|
|
13
|
+
|
|
14
|
+
class RandomOptimizer(Optimizer):
|
|
15
|
+
def __init__(self, duration:int=None, max_iterations=50):
|
|
16
|
+
"""
|
|
17
|
+
Initialize the Random Search Optimizer.
|
|
18
|
+
:param candidates: List of Candidate objects to optimize.
|
|
19
|
+
:param max_iterations: Number of random samples to evaluate.
|
|
20
|
+
:param patience: Number of iterations without improvement before stopping.
|
|
21
|
+
"""
|
|
22
|
+
self.max_iterations = max_iterations
|
|
23
|
+
self.current_iteration = 0
|
|
24
|
+
self.initial_modifier:float = 5
|
|
25
|
+
self.max_candidates = 40
|
|
26
|
+
self.first_candidates = None
|
|
27
|
+
|
|
28
|
+
|
|
29
|
+
def _randomize_hyperparameters(self, candidate):
|
|
30
|
+
"""Randomly modifies the hyperparameters of a given candidate."""
|
|
31
|
+
for _, step in candidate.pipeline.steps:
|
|
32
|
+
for key, config in step.configuration.items():
|
|
33
|
+
# Get random values
|
|
34
|
+
|
|
35
|
+
if type(config['value']) in [int, float]: # Numeric value ? Let's apply multiplier
|
|
36
|
+
is_int = isinstance(config['value'], int)
|
|
37
|
+
|
|
38
|
+
new_value = None
|
|
39
|
+
if 'range' in config: # Random in range
|
|
40
|
+
new_value = random.uniform(*config['range'])
|
|
41
|
+
else: # Strong multiplier -> kind of random
|
|
42
|
+
# Randomly choose a positive or negative editing
|
|
43
|
+
if bool(random.getrandbits(1)):
|
|
44
|
+
# Negative -> Multiply value by something between 0.01 and 1
|
|
45
|
+
change_rate = random.uniform(0.01, 1)
|
|
46
|
+
new_value = config['value']*change_rate
|
|
47
|
+
else:
|
|
48
|
+
# Positive -> Multiply value by something between 1
|
|
49
|
+
# and the max modificator in configuration
|
|
50
|
+
change_rate = random.uniform(1, self.initial_modifier)
|
|
51
|
+
new_value = config['value']*change_rate
|
|
52
|
+
|
|
53
|
+
# Value was a int ? Round it to keep it int
|
|
54
|
+
if is_int:
|
|
55
|
+
new_value = round(new_value)
|
|
56
|
+
|
|
57
|
+
if not self.__valide_config(config, new_value):
|
|
58
|
+
# Cancel is the new value is not correct.
|
|
59
|
+
new_value = config['value']
|
|
60
|
+
|
|
61
|
+
# Categorical value, choose randomly one of them
|
|
62
|
+
elif 'categorical' in config.keys():
|
|
63
|
+
new_value = random.choice(config['categorical'])
|
|
64
|
+
elif isinstance(config['value'], bool):
|
|
65
|
+
# Bool value, choose randomly beetwen True and False
|
|
66
|
+
new_value = random.choice([True, False])
|
|
67
|
+
else: # Other value ? Just keep it
|
|
68
|
+
new_value = config['value']
|
|
69
|
+
|
|
70
|
+
step.configure(key, new_value)
|
|
71
|
+
|
|
72
|
+
return candidate
|
|
73
|
+
|
|
74
|
+
def run(self, candidates):
|
|
75
|
+
"""
|
|
76
|
+
Perform a single iteration of random search optimization.
|
|
77
|
+
:param evaluate_fn: Function that takes a Candidate object and returns a performance score.
|
|
78
|
+
"""
|
|
79
|
+
candidates = candidates[0:self.max_candidates//2]
|
|
80
|
+
new_candidates = []
|
|
81
|
+
for candidate in candidates:
|
|
82
|
+
new_candidates.append(self._randomize_hyperparameters(deepcopy(candidate)))
|
|
83
|
+
|
|
84
|
+
self.current_iteration += 1
|
|
85
|
+
return candidates + new_candidates
|
|
86
|
+
|
|
87
|
+
@property
|
|
88
|
+
def finished(self) -> bool:
|
|
89
|
+
"""
|
|
90
|
+
Does optimisation is finished ?
|
|
91
|
+
|
|
92
|
+
Returns:
|
|
93
|
+
bool: finished ?
|
|
94
|
+
"""
|
|
95
|
+
return self.current_iteration >= self.max_iterations
|
|
96
|
+
|
|
97
|
+
# Does the configuration is valid or not ?
|
|
98
|
+
def __valide_config(self, config:dict, value:any) -> bool:
|
|
99
|
+
if 'range' not in config.keys():
|
|
100
|
+
return True
|
|
101
|
+
return config['range'][0] <= value <= config['range'][1]
|
iaml/plot.py
ADDED
|
@@ -0,0 +1,138 @@
|
|
|
1
|
+
"""[PLOT] Parent of all others Plot, implement the default behavior"""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import base64
|
|
5
|
+
from functools import wraps
|
|
6
|
+
from typing import Any, TYPE_CHECKING
|
|
7
|
+
|
|
8
|
+
import io
|
|
9
|
+
import matplotlib.pyplot as plt
|
|
10
|
+
import pandas as pd
|
|
11
|
+
|
|
12
|
+
if TYPE_CHECKING:
|
|
13
|
+
from .iaml_pipeline import IAMLPipeline
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
def capture(func) -> Any:
|
|
17
|
+
"""A decorator to capture a matplotlib plot into a BytesIO object and return it
|
|
18
|
+
as binary data instead of showing it.
|
|
19
|
+
|
|
20
|
+
:return: binary plot data.
|
|
21
|
+
"""
|
|
22
|
+
@wraps(func)
|
|
23
|
+
def wrapper(*args, **kwargs):
|
|
24
|
+
ret = func(*args, **kwargs)
|
|
25
|
+
plt.close()
|
|
26
|
+
return ret
|
|
27
|
+
|
|
28
|
+
return wrapper
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
class Plot:
|
|
32
|
+
"""[PLOT] Parent of all others Plot, implement the default behavior"""
|
|
33
|
+
|
|
34
|
+
title: str = "Here is the plot title"
|
|
35
|
+
"""Plot title"""
|
|
36
|
+
|
|
37
|
+
description: str = "Here is an explanation of how this plot work"
|
|
38
|
+
"""Plot description"""
|
|
39
|
+
|
|
40
|
+
def __init__(self, *_):
|
|
41
|
+
self._binary_image: io.BytesIO = None
|
|
42
|
+
"""The generated plot image"""
|
|
43
|
+
|
|
44
|
+
@property
|
|
45
|
+
def image(self) -> bytes:
|
|
46
|
+
"""Get binary representation of the plot image
|
|
47
|
+
|
|
48
|
+
:raise AttributeError: Plot must be computed before.
|
|
49
|
+
:return: The image bytes.
|
|
50
|
+
"""
|
|
51
|
+
if self._binary_image is not None:
|
|
52
|
+
return self._binary_image.getvalue()
|
|
53
|
+
|
|
54
|
+
raise AttributeError("Plot must be computed before")
|
|
55
|
+
|
|
56
|
+
@property
|
|
57
|
+
def b64_image(self) -> str:
|
|
58
|
+
"""Get b64 representation of the plot image
|
|
59
|
+
|
|
60
|
+
:return: b64 string image.
|
|
61
|
+
"""
|
|
62
|
+
return base64.b64encode(self.image).decode()
|
|
63
|
+
|
|
64
|
+
def compute(
|
|
65
|
+
self,
|
|
66
|
+
estimator: IAMLPipeline,
|
|
67
|
+
X: pd.DataFrame,
|
|
68
|
+
y: pd.Series,
|
|
69
|
+
**kwargs) -> None:
|
|
70
|
+
"""Compute plot given X, y.
|
|
71
|
+
Must be overloaded by children classes
|
|
72
|
+
|
|
73
|
+
:param IAMLPipeline estimator: The pipeline we compute the plot on.
|
|
74
|
+
:param pd.DataFrame X: The dataset we wanna compute plot on.
|
|
75
|
+
:param pd.Series y: The dataset target we wanna compute plot on.
|
|
76
|
+
:param optional \\**kwargs: Additional parameters for plotting.
|
|
77
|
+
"""
|
|
78
|
+
raise NotImplementedError('Subclass must implement abstract method')
|
|
79
|
+
|
|
80
|
+
@classmethod
|
|
81
|
+
def suitable(cls, type_of_target: str) -> bool: # pylint: disable=unused-argument
|
|
82
|
+
"""
|
|
83
|
+
Evaluates whether the plot is relevant to the type of target to
|
|
84
|
+
predict.
|
|
85
|
+
|
|
86
|
+
:param str type_of_target: Type of target to predict.
|
|
87
|
+
:return: Whether the plot is relevant.
|
|
88
|
+
"""
|
|
89
|
+
return False
|
|
90
|
+
|
|
91
|
+
def to_json(self, data_format: str='binary') -> dict[str, Any]:
|
|
92
|
+
"""Convert plot into json with name, description and b64 image
|
|
93
|
+
|
|
94
|
+
:param str data_format: Type of data we want. Can be 'binary' or 'b64'.
|
|
95
|
+
:raise AttributeError: Invalid data format.
|
|
96
|
+
:return: A dictonnary containing plot informations and data.
|
|
97
|
+
"""
|
|
98
|
+
match data_format:
|
|
99
|
+
case 'binary':
|
|
100
|
+
data = self.image
|
|
101
|
+
case 'b64':
|
|
102
|
+
data = self.b64_image
|
|
103
|
+
case _:
|
|
104
|
+
raise AttributeError('Invalid data format')
|
|
105
|
+
|
|
106
|
+
return {'title': self.title,
|
|
107
|
+
'description': self.description,
|
|
108
|
+
'image': data}
|
|
109
|
+
|
|
110
|
+
def to_markdown(self) -> str:
|
|
111
|
+
"""Return plot as markdown format
|
|
112
|
+
|
|
113
|
+
:return: Markdown formatted plot.
|
|
114
|
+
"""
|
|
115
|
+
return "\n\n".join([
|
|
116
|
+
f"# {self.title}",
|
|
117
|
+
self.description,
|
|
118
|
+
f""
|
|
119
|
+
])
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
class StatisticPlot(Plot):
|
|
123
|
+
"""Base class for descriptive statistics plots."""
|
|
124
|
+
|
|
125
|
+
enabled: bool = True
|
|
126
|
+
"""Whether this plot should be considered for rendering."""
|
|
127
|
+
|
|
128
|
+
group_by_feature: bool = False
|
|
129
|
+
"""Whether to build one plot per feature column."""
|
|
130
|
+
|
|
131
|
+
def compute(self, dataframe: pd.DataFrame, **kwargs) -> 'StatisticPlot':
|
|
132
|
+
"""Compute plot given a descriptive statistics dataframe.
|
|
133
|
+
|
|
134
|
+
:param pd.DataFrame dataframe: The descriptive statistics dataframe.
|
|
135
|
+
:param optional \\**kwargs: Additional parameters for plotting.
|
|
136
|
+
:return: A StatisticPlot object.
|
|
137
|
+
"""
|
|
138
|
+
raise NotImplementedError('Subclass must implement abstract method')
|
iaml/plots/__init__.py
ADDED
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
"""All plots"""
|
|
2
|
+
|
|
3
|
+
# Classifier
|
|
4
|
+
from .class_prediction_error_plot import ClassPredictionErrorPlot
|
|
5
|
+
from .classification_report_plot import ClassificationReportPlot
|
|
6
|
+
from .confusion_matrix_plot import ConfusionMatrixPlot
|
|
7
|
+
from .rocauc_plot import ROCAUCPlot
|
|
8
|
+
from .precision_recall_curve_plot import PrecisionRecallCurvePlot
|
|
9
|
+
|
|
10
|
+
# Regressor
|
|
11
|
+
from .residual_plot import ResidualsPlot
|
|
12
|
+
from .prediction_error_plot import PredictionErrorPlot
|
|
13
|
+
|
|
14
|
+
# Survival
|
|
15
|
+
from .kaplan_meier_comparison_plot import KaplanMeierModelComparisonPlot
|
|
16
|
+
from .cumulative_hazard_plot import CumulativeHazardModelComparisonPlot
|
|
17
|
+
from .roc_dynamique_curve_plot import ROCDynamiqueCurvePlot
|
|
18
|
+
from .shap_plot import ShapPlot
|
|
19
|
+
|
|
20
|
+
# Descriptive statistics
|
|
21
|
+
from .bar_plot import BarPlot
|
|
22
|
+
from .line_plot import LinePlot
|
|
23
|
+
from .histogram_plot import HistogramPlot
|
|
24
|
+
from .box_plot import BoxPlot
|
|
25
|
+
from .violin_plot import ViolinPlot
|
|
26
|
+
from .density_plot import DensityPlot
|
|
27
|
+
from .qq_plot import QQPlot
|
|
28
|
+
from .correlation_heatmap_plot import CorrelationHeatmapPlot
|
|
29
|
+
from .missingness_heatmap_plot import MissingnessHeatmapPlot
|
|
30
|
+
from .pair_plot import PairPlot
|
|
31
|
+
from .target_distribution_plot import TargetDistributionPlot
|
|
32
|
+
from .outlier_plot import OutlierPlot
|
iaml/plots/bar_plot.py
ADDED
|
@@ -0,0 +1,141 @@
|
|
|
1
|
+
"""[PLOT] Bar plot for descriptive statistics."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
import io
|
|
5
|
+
import textwrap
|
|
6
|
+
import numpy as np
|
|
7
|
+
import pandas as pd
|
|
8
|
+
import matplotlib.pyplot as plt
|
|
9
|
+
|
|
10
|
+
from ..plot import StatisticPlot, capture
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
def _is_missing(value: object) -> bool:
|
|
14
|
+
if value is None:
|
|
15
|
+
return True
|
|
16
|
+
if isinstance(value, float) and pd.isna(value):
|
|
17
|
+
return True
|
|
18
|
+
return False
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
def _infer_base_name(columns: list[str]) -> str:
|
|
22
|
+
for col in columns:
|
|
23
|
+
if isinstance(col, str) and col.endswith('_all'):
|
|
24
|
+
return col[:-4]
|
|
25
|
+
first = columns[0] if columns else ''
|
|
26
|
+
first_str = str(first)
|
|
27
|
+
return first_str.rsplit('_', 1)[0] if '_' in first_str else first_str
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
def _column_label(column: str, base_name: str | None) -> str:
|
|
31
|
+
column_str = str(column)
|
|
32
|
+
base_name_str = str(base_name) if base_name is not None else None
|
|
33
|
+
if base_name_str:
|
|
34
|
+
if column_str == base_name_str:
|
|
35
|
+
return 'all'
|
|
36
|
+
prefix = f"{base_name_str}_"
|
|
37
|
+
if column_str.startswith(prefix):
|
|
38
|
+
return column_str[len(prefix):]
|
|
39
|
+
if column_str.endswith('_all'):
|
|
40
|
+
return 'all'
|
|
41
|
+
return column_str.rsplit('_', 1)[-1] if '_' in column_str else 'all'
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
def _plot_placeholder(message: str) -> None:
|
|
45
|
+
plt.figure()
|
|
46
|
+
plt.text(0.5, 0.5, message, ha='center', va='center')
|
|
47
|
+
plt.axis('off')
|
|
48
|
+
|
|
49
|
+
|
|
50
|
+
class BarPlot(StatisticPlot):
|
|
51
|
+
"""[PLOT] Bar Plot."""
|
|
52
|
+
|
|
53
|
+
title: str = "Bar plot"
|
|
54
|
+
description: str = textwrap.dedent("""\
|
|
55
|
+
The bar plot shows descriptive statistics for categorical or numerical columns.
|
|
56
|
+
Categorical plots show value counts per class, while numerical plots show metrics
|
|
57
|
+
such as mean or variance per class.
|
|
58
|
+
""")
|
|
59
|
+
group_by_feature: bool = True
|
|
60
|
+
|
|
61
|
+
@capture
|
|
62
|
+
def compute(
|
|
63
|
+
self,
|
|
64
|
+
dataframe: pd.DataFrame,
|
|
65
|
+
base_name: str | None = None,
|
|
66
|
+
**kwargs) -> 'BarPlot':
|
|
67
|
+
"""Compute bar plot statistics."""
|
|
68
|
+
self._binary_image = io.BytesIO()
|
|
69
|
+
|
|
70
|
+
if dataframe.empty:
|
|
71
|
+
_plot_placeholder("No statistics available")
|
|
72
|
+
plt.savefig(self._binary_image, format='png')
|
|
73
|
+
return self
|
|
74
|
+
|
|
75
|
+
base_name = base_name or _infer_base_name(list(dataframe.columns))
|
|
76
|
+
|
|
77
|
+
if 'value_counts' in dataframe.index and dataframe.loc['value_counts'].notna().any():
|
|
78
|
+
categories = []
|
|
79
|
+
seen = set()
|
|
80
|
+
for col in dataframe.columns:
|
|
81
|
+
values = dataframe.at['value_counts', col]
|
|
82
|
+
if _is_missing(values):
|
|
83
|
+
continue
|
|
84
|
+
for cat, _ in values:
|
|
85
|
+
if cat not in seen:
|
|
86
|
+
seen.add(cat)
|
|
87
|
+
categories.append(cat)
|
|
88
|
+
|
|
89
|
+
if not categories:
|
|
90
|
+
_plot_placeholder("No categorical statistics available")
|
|
91
|
+
plt.savefig(self._binary_image, format='png')
|
|
92
|
+
return self
|
|
93
|
+
|
|
94
|
+
labels = [_column_label(col, base_name) for col in dataframe.columns]
|
|
95
|
+
data = {cat: [] for cat in categories}
|
|
96
|
+
for col in dataframe.columns:
|
|
97
|
+
values = dataframe.at['value_counts', col]
|
|
98
|
+
value_dict = dict(values) if not _is_missing(values) else {}
|
|
99
|
+
for cat in categories:
|
|
100
|
+
data[cat].append(value_dict.get(cat, 0))
|
|
101
|
+
df = pd.DataFrame(data, index=labels)
|
|
102
|
+
|
|
103
|
+
plt.figure()
|
|
104
|
+
x = np.arange(len(df.index))
|
|
105
|
+
width = min(0.8 / max(len(df.columns), 1), 0.2)
|
|
106
|
+
for i, category in enumerate(df.columns):
|
|
107
|
+
offset = (i - (len(df.columns) - 1) / 2) * width
|
|
108
|
+
plt.bar(x + offset, df[category], width, label=category)
|
|
109
|
+
plt.xticks(x, labels)
|
|
110
|
+
plt.legend()
|
|
111
|
+
plt.ylabel('Count')
|
|
112
|
+
plt.title(f'[CAT] {base_name} Statistics')
|
|
113
|
+
plt.tight_layout()
|
|
114
|
+
elif 'mean' in dataframe.index and dataframe.loc['mean'].notna().any():
|
|
115
|
+
numeric_df = dataframe.copy()
|
|
116
|
+
for row in ['mode', 'value_counts', 'null_count', 'count']:
|
|
117
|
+
if row in numeric_df.index:
|
|
118
|
+
numeric_df = numeric_df.drop(index=row)
|
|
119
|
+
numeric_df = numeric_df.dropna(how='all')
|
|
120
|
+
if numeric_df.empty:
|
|
121
|
+
_plot_placeholder("No numeric statistics available")
|
|
122
|
+
plt.savefig(self._binary_image, format='png')
|
|
123
|
+
return self
|
|
124
|
+
|
|
125
|
+
labels = [_column_label(col, base_name) for col in numeric_df.columns]
|
|
126
|
+
x = np.arange(len(numeric_df.index))
|
|
127
|
+
width = min(0.8 / max(len(numeric_df.columns), 1), 0.25)
|
|
128
|
+
plt.figure()
|
|
129
|
+
for i, col in enumerate(numeric_df.columns):
|
|
130
|
+
offset = (i - (len(numeric_df.columns) - 1) / 2) * width
|
|
131
|
+
plt.bar(x + offset, numeric_df[col], width, label=labels[i])
|
|
132
|
+
plt.xticks(x, numeric_df.index, rotation=45, ha='right')
|
|
133
|
+
plt.legend()
|
|
134
|
+
plt.ylabel('Value')
|
|
135
|
+
plt.title(f'[NUM] {base_name} Statistics')
|
|
136
|
+
plt.tight_layout()
|
|
137
|
+
else:
|
|
138
|
+
_plot_placeholder("No statistics available for bar plot")
|
|
139
|
+
|
|
140
|
+
plt.savefig(self._binary_image, format='png')
|
|
141
|
+
return self
|
iaml/plots/box_plot.py
ADDED
|
@@ -0,0 +1,166 @@
|
|
|
1
|
+
"""[PLOT] Box 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: 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 _infer_base_name(columns: list[str]) -> str:
|
|
26
|
+
for col in columns:
|
|
27
|
+
if isinstance(col, str) and col.endswith('_all'):
|
|
28
|
+
return col[:-4]
|
|
29
|
+
first = columns[0] if columns else ''
|
|
30
|
+
first_str = str(first)
|
|
31
|
+
return first_str.rsplit('_', 1)[0] if '_' in first_str else first_str
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
def _column_label(column: str, base_name: str | None) -> str:
|
|
35
|
+
column_str = str(column)
|
|
36
|
+
base_name_str = str(base_name) if base_name is not None else None
|
|
37
|
+
if base_name_str:
|
|
38
|
+
if column_str == base_name_str:
|
|
39
|
+
return 'all'
|
|
40
|
+
prefix = f"{base_name_str}_"
|
|
41
|
+
if column_str.startswith(prefix):
|
|
42
|
+
return column_str[len(prefix):]
|
|
43
|
+
if column_str.endswith('_all'):
|
|
44
|
+
return 'all'
|
|
45
|
+
return column_str.rsplit('_', 1)[-1] if '_' in column_str else column_str
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def _base_from_column(column: str) -> str:
|
|
49
|
+
column_str = str(column)
|
|
50
|
+
if column_str.endswith('_all'):
|
|
51
|
+
return column_str[:-4]
|
|
52
|
+
return column_str.rsplit('_', 1)[0] if '_' in column_str else column_str
|
|
53
|
+
|
|
54
|
+
|
|
55
|
+
def _plot_placeholder(message: str) -> None:
|
|
56
|
+
plt.figure()
|
|
57
|
+
plt.text(0.5, 0.5, message, ha='center', va='center')
|
|
58
|
+
plt.axis('off')
|
|
59
|
+
|
|
60
|
+
|
|
61
|
+
def _extract_box_stats(dataframe: pd.DataFrame, column: str) -> dict[str, float] | None:
|
|
62
|
+
required_rows = {
|
|
63
|
+
'min': 'whislo',
|
|
64
|
+
'max': 'whishi',
|
|
65
|
+
'quantile_0.25': 'q1',
|
|
66
|
+
'quantile_0.5': 'med',
|
|
67
|
+
'quantile_0.75': 'q3',
|
|
68
|
+
}
|
|
69
|
+
stats: dict[str, float] = {}
|
|
70
|
+
for row, key in required_rows.items():
|
|
71
|
+
if row not in dataframe.index:
|
|
72
|
+
return None
|
|
73
|
+
value = dataframe.at[row, column]
|
|
74
|
+
if _is_missing(value):
|
|
75
|
+
return None
|
|
76
|
+
try:
|
|
77
|
+
value = float(value)
|
|
78
|
+
except (TypeError, ValueError):
|
|
79
|
+
return None
|
|
80
|
+
if not np.isfinite(value):
|
|
81
|
+
return None
|
|
82
|
+
stats[key] = value
|
|
83
|
+
stats['fliers'] = []
|
|
84
|
+
return stats
|
|
85
|
+
|
|
86
|
+
|
|
87
|
+
class BoxPlot(StatisticPlot):
|
|
88
|
+
"""[PLOT] Box Plot."""
|
|
89
|
+
|
|
90
|
+
name: str = "Box Plot"
|
|
91
|
+
_description: str = textwrap.dedent("""\
|
|
92
|
+
Box plots summarize numeric distributions per class.
|
|
93
|
+
""")
|
|
94
|
+
_description_long: str = textwrap.dedent("""\
|
|
95
|
+
Box plots visualize min, quartiles, and max values for numeric columns
|
|
96
|
+
per class, using precomputed descriptive statistics.
|
|
97
|
+
""")
|
|
98
|
+
refs: list[dict] = []
|
|
99
|
+
|
|
100
|
+
title: str = "Box plot"
|
|
101
|
+
description: str = textwrap.dedent("""\
|
|
102
|
+
The box plot shows numerical distributions per class.
|
|
103
|
+
""")
|
|
104
|
+
group_by_feature: bool = True
|
|
105
|
+
|
|
106
|
+
def __str__(self) -> str:
|
|
107
|
+
return 'boxplot'
|
|
108
|
+
|
|
109
|
+
@capture
|
|
110
|
+
def compute(
|
|
111
|
+
self,
|
|
112
|
+
dataframe: pd.DataFrame,
|
|
113
|
+
base_name: str | None = None,
|
|
114
|
+
dataset: Dataset | None = None,
|
|
115
|
+
**kwargs) -> 'BoxPlot':
|
|
116
|
+
"""Compute box plot statistics."""
|
|
117
|
+
self._binary_image = io.BytesIO()
|
|
118
|
+
|
|
119
|
+
if dataframe.empty:
|
|
120
|
+
_plot_placeholder("No statistics available")
|
|
121
|
+
plt.savefig(self._binary_image, format='png')
|
|
122
|
+
return self
|
|
123
|
+
|
|
124
|
+
base_name = base_name or _infer_base_name(list(dataframe.columns))
|
|
125
|
+
|
|
126
|
+
if dataset is not None:
|
|
127
|
+
numeric_columns = set(dataset.get_columns_names_by_type(DataType.NUMERIC))
|
|
128
|
+
if base_name in numeric_columns:
|
|
129
|
+
columns_to_show = [
|
|
130
|
+
col for col in dataframe.columns
|
|
131
|
+
if _base_from_column(col) == base_name
|
|
132
|
+
]
|
|
133
|
+
else:
|
|
134
|
+
columns_to_show = [
|
|
135
|
+
col for col in dataframe.columns
|
|
136
|
+
if _base_from_column(col) in numeric_columns
|
|
137
|
+
]
|
|
138
|
+
else:
|
|
139
|
+
columns_to_show = list(dataframe.columns)
|
|
140
|
+
|
|
141
|
+
if not columns_to_show:
|
|
142
|
+
_plot_placeholder("No numeric statistics available")
|
|
143
|
+
plt.savefig(self._binary_image, format='png')
|
|
144
|
+
return self
|
|
145
|
+
|
|
146
|
+
entries: list[dict[str, Any]] = []
|
|
147
|
+
for col in columns_to_show:
|
|
148
|
+
stats = _extract_box_stats(dataframe, col)
|
|
149
|
+
if stats is None:
|
|
150
|
+
continue
|
|
151
|
+
stats['label'] = _column_label(col, base_name)
|
|
152
|
+
entries.append(stats)
|
|
153
|
+
|
|
154
|
+
if not entries:
|
|
155
|
+
_plot_placeholder("No box plot statistics available")
|
|
156
|
+
plt.savefig(self._binary_image, format='png')
|
|
157
|
+
return self
|
|
158
|
+
|
|
159
|
+
plt.figure(figsize=(max(4.0, 0.9 * len(entries)), 4.0))
|
|
160
|
+
plt.bxp(entries, showfliers=False)
|
|
161
|
+
plt.ylabel('Value')
|
|
162
|
+
if base_name:
|
|
163
|
+
plt.title(f"Box plot: {base_name}")
|
|
164
|
+
plt.tight_layout()
|
|
165
|
+
plt.savefig(self._binary_image, format='png')
|
|
166
|
+
return self
|
|
@@ -0,0 +1,37 @@
|
|
|
1
|
+
"""[PLOT] Class Prediction Error Plot"""
|
|
2
|
+
import textwrap
|
|
3
|
+
|
|
4
|
+
from yellowbrick.classifier import ClassPredictionError
|
|
5
|
+
|
|
6
|
+
from ..metric_plot import MetricPlot, yellowbrick_plot
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@yellowbrick_plot(ClassPredictionError)
|
|
10
|
+
class ClassPredictionErrorPlot(MetricPlot):
|
|
11
|
+
"""[PLOT] Class Prediction Error Plot"""
|
|
12
|
+
|
|
13
|
+
title: str = "Prediction Error Plot"
|
|
14
|
+
description: str = textwrap.dedent("""
|
|
15
|
+
The Class Prediction Error is a visualization that helps understand how well a machine
|
|
16
|
+
learning model is performing in predicting medical conditions or diagnoses. It shows
|
|
17
|
+
both the correct predictions made by the model and the mistakes it makes for each condition.
|
|
18
|
+
|
|
19
|
+
For example, imagine you have a model that’s trained to identify different diseases
|
|
20
|
+
from patient data, such as predicting whether someone has diabetes, hypertension,
|
|
21
|
+
or is healthy. The Class Prediction Error plot would show, for each of these conditions,
|
|
22
|
+
how many times the model correctly identified the disease and how many times it made a
|
|
23
|
+
wrong prediction.
|
|
24
|
+
|
|
25
|
+
For instance, if the model predicts "diabetes" for a patient who actually has "hypertension,"
|
|
26
|
+
the plot will highlight this error. Similarly, it will also show how often the model
|
|
27
|
+
correctly identifies "healthy" patients versus when it mistakenly predicts they have
|
|
28
|
+
a disease.
|
|
29
|
+
|
|
30
|
+
This visualization is especially helpful for doctors and data scientists because it clearly
|
|
31
|
+
shows where the model is making errors, making it easier to improve its accuracy, which is
|
|
32
|
+
critical in healthcare where correct predictions can have a big impact on patient outcomes.
|
|
33
|
+
""")
|
|
34
|
+
|
|
35
|
+
@classmethod
|
|
36
|
+
def suitable(cls, type_of_target: str) -> bool:
|
|
37
|
+
return type_of_target in ['binary', 'multiclass', 'multilabel-indicator']
|
|
@@ -0,0 +1,35 @@
|
|
|
1
|
+
"""[PLOT] Classification Report Plot"""
|
|
2
|
+
import textwrap
|
|
3
|
+
|
|
4
|
+
from yellowbrick.classifier import ClassificationReport
|
|
5
|
+
|
|
6
|
+
from ..metric_plot import MetricPlot, yellowbrick_plot
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@yellowbrick_plot(ClassificationReport)
|
|
10
|
+
class ClassificationReportPlot(MetricPlot):
|
|
11
|
+
"""[PLOT] Classification Report Plot"""
|
|
12
|
+
|
|
13
|
+
title: str = "Classification report"
|
|
14
|
+
description: str = textwrap.dedent("""
|
|
15
|
+
The Classification Report is a visual tool to evaluate the performance of a machine learning model
|
|
16
|
+
on classification tasks, such as diagnosing medical conditions. This plot provides key metrics for
|
|
17
|
+
each class (like diseases or health conditions) that the model is trained to identify.
|
|
18
|
+
|
|
19
|
+
It includes metrics such as precision, recall, F1-score, and support for each class. These metrics
|
|
20
|
+
are essential to understanding how well the model is identifying true positives (correct diagnoses),
|
|
21
|
+
minimizing false positives (incorrect diagnoses), and balancing between precision and recall.
|
|
22
|
+
|
|
23
|
+
For example, if you have a model classifying conditions like 'healthy', 'diabetes', and 'hypertension',
|
|
24
|
+
the Classification Report plot will show you the precision (how many of the predicted conditions were correct),
|
|
25
|
+
recall (how many actual conditions were correctly identified), and the F1-score (the harmonic mean of precision
|
|
26
|
+
and recall). This is especially important in healthcare to ensure the model provides balanced and accurate results
|
|
27
|
+
across all classes, improving both diagnosis and patient outcomes.
|
|
28
|
+
|
|
29
|
+
Doctors and data scientists use this visualization to easily compare the model's performance on different conditions,
|
|
30
|
+
aiding in model refinement and ensuring robust diagnostic predictions.
|
|
31
|
+
""")
|
|
32
|
+
|
|
33
|
+
@classmethod
|
|
34
|
+
def suitable(cls, type_of_target: str) -> bool:
|
|
35
|
+
return type_of_target in ['binary', 'multiclass', 'multilabel-indicator']
|
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
"""[PLOT] Confusion Matrix Plot"""
|
|
2
|
+
import textwrap
|
|
3
|
+
|
|
4
|
+
from yellowbrick.classifier import ConfusionMatrix
|
|
5
|
+
|
|
6
|
+
from ..metric_plot import MetricPlot, yellowbrick_plot
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
@yellowbrick_plot(ConfusionMatrix)
|
|
10
|
+
class ConfusionMatrixPlot(MetricPlot):
|
|
11
|
+
"""[PLOT] Confusion Matrix Plot"""
|
|
12
|
+
|
|
13
|
+
title: str = "Confusion Matrix"
|
|
14
|
+
description: str = textwrap.dedent("""
|
|
15
|
+
The Confusion Matrix is a powerful visualization tool used to assess how well a classification model
|
|
16
|
+
is performing, particularly in identifying different medical conditions or diagnostic categories.
|
|
17
|
+
It provides a clear breakdown of true positive, false positive, true negative, and false negative rates.
|
|
18
|
+
|
|
19
|
+
For example, in a healthcare setting, if a model is trained to classify whether a patient has 'diabetes',
|
|
20
|
+
'hypertension', or is 'healthy', the Confusion Matrix will show how many times the model made correct predictions
|
|
21
|
+
(true positives and true negatives) and where it made mistakes (false positives and false negatives).
|
|
22
|
+
|
|
23
|
+
Each row in the matrix represents the actual condition of the patient, and each column represents the predicted
|
|
24
|
+
condition. This makes it easy to spot patterns in the model's predictions, such as whether it tends to misclassify
|
|
25
|
+
one condition as another.
|
|
26
|
+
|
|
27
|
+
Doctors and data scientists rely on this plot to understand not only how often the model is right but also the
|
|
28
|
+
types of errors it makes. This information is crucial in healthcare, where reducing misdiagnoses can significantly
|
|
29
|
+
improve patient outcomes.
|
|
30
|
+
""")
|
|
31
|
+
|
|
32
|
+
@classmethod
|
|
33
|
+
def suitable(cls, type_of_target: str) -> bool:
|
|
34
|
+
return type_of_target in ['binary', 'multiclass', 'multilabel-indicator']
|