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,206 @@
|
|
|
1
|
+
"""[PLOT] Violin 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 _column_label(column: str, base_name: str | None) -> str:
|
|
26
|
+
column_str = str(column)
|
|
27
|
+
base_name_str = str(base_name) if base_name is not None else None
|
|
28
|
+
if base_name_str:
|
|
29
|
+
if column_str == base_name_str:
|
|
30
|
+
return 'all'
|
|
31
|
+
prefix = f"{base_name_str}_"
|
|
32
|
+
if column_str.startswith(prefix):
|
|
33
|
+
return column_str[len(prefix):]
|
|
34
|
+
if column_str.endswith('_all'):
|
|
35
|
+
return 'all'
|
|
36
|
+
return column_str
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
def _plot_placeholder(message: str) -> None:
|
|
40
|
+
plt.figure()
|
|
41
|
+
plt.text(0.5, 0.5, message, ha='center', va='center')
|
|
42
|
+
plt.axis('off')
|
|
43
|
+
|
|
44
|
+
|
|
45
|
+
def _extract_violin_groups(
|
|
46
|
+
value: Any,
|
|
47
|
+
) -> list[tuple[str, np.ndarray, np.ndarray, np.ndarray | None]]:
|
|
48
|
+
if _is_missing(value) or not isinstance(value, dict):
|
|
49
|
+
return []
|
|
50
|
+
entries: list[tuple[str, np.ndarray, np.ndarray, np.ndarray | None]] = []
|
|
51
|
+
for category, stats in value.items():
|
|
52
|
+
if not isinstance(stats, dict):
|
|
53
|
+
continue
|
|
54
|
+
density = stats.get('density')
|
|
55
|
+
support = stats.get('support')
|
|
56
|
+
if density is None or support is None:
|
|
57
|
+
continue
|
|
58
|
+
try:
|
|
59
|
+
density_arr = np.asarray(density, dtype=float).ravel()
|
|
60
|
+
support_arr = np.asarray(support, dtype=float).ravel()
|
|
61
|
+
except (TypeError, ValueError):
|
|
62
|
+
continue
|
|
63
|
+
if density_arr.size == 0 or support_arr.size == 0:
|
|
64
|
+
continue
|
|
65
|
+
if density_arr.size != support_arr.size:
|
|
66
|
+
continue
|
|
67
|
+
mask = np.isfinite(density_arr) & np.isfinite(support_arr)
|
|
68
|
+
if not np.any(mask):
|
|
69
|
+
continue
|
|
70
|
+
density_arr = density_arr[mask]
|
|
71
|
+
support_arr = support_arr[mask]
|
|
72
|
+
order = np.argsort(support_arr)
|
|
73
|
+
density_arr = density_arr[order]
|
|
74
|
+
support_arr = support_arr[order]
|
|
75
|
+
|
|
76
|
+
quartiles = None
|
|
77
|
+
if 'quartiles' in stats:
|
|
78
|
+
try:
|
|
79
|
+
q_values = np.asarray(stats['quartiles'], dtype=float).ravel()
|
|
80
|
+
except (TypeError, ValueError):
|
|
81
|
+
q_values = np.asarray([], dtype=float)
|
|
82
|
+
if q_values.size >= 3 and np.all(np.isfinite(q_values[:3])):
|
|
83
|
+
quartiles = q_values[:3]
|
|
84
|
+
entries.append((str(category), support_arr, density_arr, quartiles))
|
|
85
|
+
return entries
|
|
86
|
+
|
|
87
|
+
|
|
88
|
+
def _render_violin(
|
|
89
|
+
ax: plt.Axes,
|
|
90
|
+
groups: list[tuple[str, np.ndarray, np.ndarray, np.ndarray | None]],
|
|
91
|
+
) -> None:
|
|
92
|
+
if not groups:
|
|
93
|
+
ax.text(0.5, 0.5, "No violin data", ha='center', va='center')
|
|
94
|
+
ax.axis('off')
|
|
95
|
+
return
|
|
96
|
+
max_density = max(float(np.nanmax(density)) for _, _, density, _ in groups)
|
|
97
|
+
if not np.isfinite(max_density) or max_density <= 0:
|
|
98
|
+
ax.text(0.5, 0.5, "Invalid density data", ha='center', va='center')
|
|
99
|
+
ax.axis('off')
|
|
100
|
+
return
|
|
101
|
+
|
|
102
|
+
scale = 0.4 / max_density
|
|
103
|
+
for idx, (category, support, density, quartiles) in enumerate(groups):
|
|
104
|
+
width = density * scale
|
|
105
|
+
ax.fill_betweenx(
|
|
106
|
+
support,
|
|
107
|
+
idx - width,
|
|
108
|
+
idx + width,
|
|
109
|
+
alpha=0.6,
|
|
110
|
+
edgecolor='black',
|
|
111
|
+
linewidth=0.8,
|
|
112
|
+
)
|
|
113
|
+
if quartiles is not None:
|
|
114
|
+
q1, median, q3 = quartiles
|
|
115
|
+
ax.plot([idx, idx], [q1, q3], color='black', linewidth=2)
|
|
116
|
+
ax.plot([idx - 0.08, idx + 0.08], [median, median], color='black', linewidth=2)
|
|
117
|
+
|
|
118
|
+
ax.set_xticks(range(len(groups)))
|
|
119
|
+
ax.set_xticklabels([label for label, *_ in groups], rotation=30, ha='right')
|
|
120
|
+
ax.set_xlabel('Category')
|
|
121
|
+
ax.set_ylabel('Target')
|
|
122
|
+
|
|
123
|
+
|
|
124
|
+
class ViolinPlot(StatisticPlot):
|
|
125
|
+
"""[PLOT] Violin Plot."""
|
|
126
|
+
|
|
127
|
+
name: str = "Violin Plot"
|
|
128
|
+
_description: str = textwrap.dedent("""\
|
|
129
|
+
Violin plots show target distributions per categorical value.
|
|
130
|
+
""")
|
|
131
|
+
_description_long: str = textwrap.dedent("""\
|
|
132
|
+
This plot renders violin distributions for categorical features using
|
|
133
|
+
precomputed density statistics for a continuous target.
|
|
134
|
+
""")
|
|
135
|
+
refs: list[dict] = []
|
|
136
|
+
|
|
137
|
+
title: str = "Violin plot"
|
|
138
|
+
description: str = textwrap.dedent("""\
|
|
139
|
+
The violin plot shows target distributions per categorical feature.
|
|
140
|
+
""")
|
|
141
|
+
group_by_feature: bool = True
|
|
142
|
+
|
|
143
|
+
def __str__(self) -> str:
|
|
144
|
+
return 'violin'
|
|
145
|
+
|
|
146
|
+
@capture
|
|
147
|
+
def compute(
|
|
148
|
+
self,
|
|
149
|
+
dataframe: pd.DataFrame,
|
|
150
|
+
base_name: str | None = None,
|
|
151
|
+
dataset: Dataset | None = None,
|
|
152
|
+
**kwargs,
|
|
153
|
+
) -> 'ViolinPlot':
|
|
154
|
+
"""Compute violin plot statistics."""
|
|
155
|
+
self._binary_image = io.BytesIO()
|
|
156
|
+
|
|
157
|
+
if dataframe.empty:
|
|
158
|
+
_plot_placeholder("No statistics available")
|
|
159
|
+
plt.savefig(self._binary_image, format='png')
|
|
160
|
+
return self
|
|
161
|
+
|
|
162
|
+
violin_key = str(self)
|
|
163
|
+
if violin_key not in dataframe.index:
|
|
164
|
+
_plot_placeholder("Violin statistics not available")
|
|
165
|
+
plt.savefig(self._binary_image, format='png')
|
|
166
|
+
return self
|
|
167
|
+
|
|
168
|
+
if dataset is not None:
|
|
169
|
+
categorical_columns = set(dataset.get_columns_names_by_type(DataType.CATEGORICAL))
|
|
170
|
+
columns_to_show = [col for col in dataframe.columns if col in categorical_columns]
|
|
171
|
+
else:
|
|
172
|
+
columns_to_show = list(dataframe.columns)
|
|
173
|
+
|
|
174
|
+
violin_row = dataframe.loc[violin_key]
|
|
175
|
+
entries: list[tuple[str, list[tuple[str, np.ndarray, np.ndarray, np.ndarray | None]]]] = []
|
|
176
|
+
for col in columns_to_show:
|
|
177
|
+
groups = _extract_violin_groups(violin_row.get(col))
|
|
178
|
+
if groups:
|
|
179
|
+
entries.append((col, groups))
|
|
180
|
+
|
|
181
|
+
if not entries:
|
|
182
|
+
_plot_placeholder("No violin statistics available")
|
|
183
|
+
plt.savefig(self._binary_image, format='png')
|
|
184
|
+
return self
|
|
185
|
+
|
|
186
|
+
n_plots = len(entries)
|
|
187
|
+
n_cols = 1 if n_plots == 1 else 2
|
|
188
|
+
n_rows = int(np.ceil(n_plots / n_cols))
|
|
189
|
+
fig, axes = plt.subplots(n_rows, n_cols, figsize=(6 * n_cols, 4 * n_rows))
|
|
190
|
+
axes_list = np.atleast_1d(axes).ravel()
|
|
191
|
+
|
|
192
|
+
for ax, (col, groups) in zip(axes_list, entries):
|
|
193
|
+
_render_violin(ax, groups)
|
|
194
|
+
ax.set_title(_column_label(col, base_name))
|
|
195
|
+
|
|
196
|
+
for ax in axes_list[len(entries):]:
|
|
197
|
+
ax.axis('off')
|
|
198
|
+
|
|
199
|
+
if base_name:
|
|
200
|
+
fig.suptitle(f"Violin plot: {base_name}")
|
|
201
|
+
plt.tight_layout(rect=(0, 0, 1, 0.95))
|
|
202
|
+
else:
|
|
203
|
+
plt.tight_layout()
|
|
204
|
+
|
|
205
|
+
plt.savefig(self._binary_image, format='png')
|
|
206
|
+
return self
|
iaml/predictor.py
ADDED
|
@@ -0,0 +1,139 @@
|
|
|
1
|
+
"""Last step of a pipeline -> can make prediction"""
|
|
2
|
+
from abc import ABCMeta, abstractmethod
|
|
3
|
+
from typing import Any
|
|
4
|
+
import dataclasses
|
|
5
|
+
import pandas as pd
|
|
6
|
+
from sklearn.base import BaseEstimator
|
|
7
|
+
from .actionable import Actionable
|
|
8
|
+
from .decorators.runner import runner
|
|
9
|
+
from .candidate import Candidate
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@dataclasses.dataclass
|
|
13
|
+
class Model(metaclass=ABCMeta):
|
|
14
|
+
"""Model type (use for typing purposes only)."""
|
|
15
|
+
|
|
16
|
+
@abstractmethod
|
|
17
|
+
def predict(self, X: Any, *args, **kw) -> Any:
|
|
18
|
+
"""Any predict method implemented by most ML frameworks.
|
|
19
|
+
|
|
20
|
+
:param Any X: The dataset to predict
|
|
21
|
+
:param tuple, optional \\*args: Additional parameters.
|
|
22
|
+
:param tuple, optional \\**kwargs: Additional parameters.
|
|
23
|
+
:return: Prediction output.
|
|
24
|
+
"""
|
|
25
|
+
|
|
26
|
+
class Predictor(Actionable, BaseEstimator, metaclass=ABCMeta):
|
|
27
|
+
"""[STEP] Abstract learning step
|
|
28
|
+
|
|
29
|
+
Also acts as an interface with traditional scikit-learn models
|
|
30
|
+
for better integration with IAMLPipeline.
|
|
31
|
+
"""
|
|
32
|
+
|
|
33
|
+
model: Model
|
|
34
|
+
|
|
35
|
+
def __init__(self):
|
|
36
|
+
super().__init__()
|
|
37
|
+
self.optimizable: bool = True
|
|
38
|
+
self.model: Model = None
|
|
39
|
+
|
|
40
|
+
@runner
|
|
41
|
+
def run(self, candidate: Candidate) -> Candidate:
|
|
42
|
+
"""Run the step. In "Run" stage, predict does not "fit". Only add himself to pipeline
|
|
43
|
+
|
|
44
|
+
:param Candidate candidate: Candidate informations
|
|
45
|
+
:return: transformed Candidate
|
|
46
|
+
"""
|
|
47
|
+
return candidate.add_to_pipeline(self)
|
|
48
|
+
|
|
49
|
+
def predict_proba(self, X: pd.DataFrame) -> list[float]:
|
|
50
|
+
"""Apply prediction model on DataFrame with probability
|
|
51
|
+
|
|
52
|
+
:param pd.DataFrame X: DataFrame use to predict
|
|
53
|
+
:raise AttributeError: Unable to predict probabilities with this model
|
|
54
|
+
:return: Predicted values
|
|
55
|
+
"""
|
|
56
|
+
if self.model and hasattr(self.model, 'predict_proba'):
|
|
57
|
+
return self.model.predict_proba(X)
|
|
58
|
+
|
|
59
|
+
raise AttributeError("Unable to predict probabilities with this model")
|
|
60
|
+
|
|
61
|
+
def __getattribute__(self, attr: str) -> bool:
|
|
62
|
+
"""Overload getattr to allow accurate hasattr on predict_proba
|
|
63
|
+
|
|
64
|
+
:param str attr: Attribute to test.
|
|
65
|
+
:raise AttributeError: predict_proba not implemented in this model.
|
|
66
|
+
:return: Is attribute implemented ?
|
|
67
|
+
"""
|
|
68
|
+
if attr == 'predict_proba' \
|
|
69
|
+
and not( \
|
|
70
|
+
self.model and hasattr(self.model, 'predict_proba') \
|
|
71
|
+
):
|
|
72
|
+
raise AttributeError("predict_proba not implemented in this model")
|
|
73
|
+
|
|
74
|
+
return super().__getattribute__(attr)
|
|
75
|
+
|
|
76
|
+
def predict(self, X: pd.DataFrame) -> list[float]:
|
|
77
|
+
"""Apply prediction model on DataFrame
|
|
78
|
+
|
|
79
|
+
:param pd.DataFrame X: DataFrame use to predict.
|
|
80
|
+
:return: Predicted values.
|
|
81
|
+
"""
|
|
82
|
+
if self.model and hasattr(self.model, 'predict'):
|
|
83
|
+
results = self.model.predict(X)
|
|
84
|
+
if hasattr(self, 'label_encoder'):
|
|
85
|
+
return self.label_encoder.inverse_transform(results)
|
|
86
|
+
return results
|
|
87
|
+
return None
|
|
88
|
+
|
|
89
|
+
def predict_survival_function(self, X: pd.DataFrame) -> list[list[float]]:
|
|
90
|
+
"""Apply prediction survival function model on DataFrame
|
|
91
|
+
|
|
92
|
+
:param pd.DataFrame X: DataFrame use to predict.
|
|
93
|
+
:raise AttributeError: Unable to predict survival function with this model.
|
|
94
|
+
:return: Predicted values.
|
|
95
|
+
"""
|
|
96
|
+
if self.model and hasattr(self.model, 'predict_survival_function'):
|
|
97
|
+
return self.model.predict_survival_function(X)
|
|
98
|
+
|
|
99
|
+
raise AttributeError("Unable to predict survival function with this model")
|
|
100
|
+
|
|
101
|
+
def predict_cumulative_hazard_function(self, X: pd.DataFrame) -> list[list[float]]:
|
|
102
|
+
"""Apply prediction survival function model on DataFrame
|
|
103
|
+
|
|
104
|
+
:param pd.DataFrame X: DataFrame use to predict.
|
|
105
|
+
:raise AttributeError: Unable to predict survival function with this model.
|
|
106
|
+
:return: Predicted values.
|
|
107
|
+
"""
|
|
108
|
+
if self.model and hasattr(self.model, 'predict_cumulative_hazard_function'):
|
|
109
|
+
return self.model.predict_cumulative_hazard_function(X)
|
|
110
|
+
|
|
111
|
+
raise AttributeError("Unable to predict cumulative hazard function with this model")
|
|
112
|
+
|
|
113
|
+
@property
|
|
114
|
+
def classes_(self) -> list:
|
|
115
|
+
"""Return classes of the target in fit data"""
|
|
116
|
+
return self.model.classes_
|
|
117
|
+
|
|
118
|
+
def score(self, *args, **kwargs) -> Any:
|
|
119
|
+
"""Mimic Scikitlearn API
|
|
120
|
+
|
|
121
|
+
:return: The model score
|
|
122
|
+
"""
|
|
123
|
+
return self.model.score(*args, **kwargs)
|
|
124
|
+
|
|
125
|
+
def get_params(self, *args, **kwargs) -> Any | None:
|
|
126
|
+
"""Mimic Scikitlearn API
|
|
127
|
+
|
|
128
|
+
:return: Model parameters or None if no parameters
|
|
129
|
+
"""
|
|
130
|
+
if self.model and hasattr(self.model, 'get_params'):
|
|
131
|
+
return self.model.get_params(*args, **kwargs)
|
|
132
|
+
return None
|
|
133
|
+
|
|
134
|
+
def __name__(self) -> str:
|
|
135
|
+
"""Return the predictor formatted name
|
|
136
|
+
|
|
137
|
+
:return: formatted name
|
|
138
|
+
"""
|
|
139
|
+
return ' '.join(x.title() for x in str(self).split('_'))
|
iaml/reference.py
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Reference class.
|
|
3
|
+
Contain all needed data to provide a reference for a step
|
|
4
|
+
"""
|
|
5
|
+
from typing import List, Dict
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class Reference: # pylint: disable=too-few-public-methods
|
|
9
|
+
"""Reference class.
|
|
10
|
+
Contain all needed data to provide a reference for a step
|
|
11
|
+
|
|
12
|
+
:param Dict properties: a dictionnary of properties for a reference object
|
|
13
|
+
:param str step_name: The step_name attached to this reference
|
|
14
|
+
"""
|
|
15
|
+
|
|
16
|
+
def __init__(self, properties: Dict, step_name: str) -> None:
|
|
17
|
+
"""Instantiate all properties provided to the specific reference
|
|
18
|
+
such as year of publication, authors, doi ...
|
|
19
|
+
"""
|
|
20
|
+
setattr(self, 'step', step_name)
|
|
21
|
+
for k, v in properties.items():
|
|
22
|
+
setattr(self, k, v)
|
|
23
|
+
|
|
24
|
+
def __str__(self) -> str:
|
|
25
|
+
"""Return a simple string containing reference information
|
|
26
|
+
|
|
27
|
+
:return: The reference representation
|
|
28
|
+
"""
|
|
29
|
+
structured = ''
|
|
30
|
+
try:
|
|
31
|
+
structured = structured + ', '.join(self.authors) + '. '
|
|
32
|
+
except AttributeError:
|
|
33
|
+
pass
|
|
34
|
+
try:
|
|
35
|
+
structured = structured + str(self.name) + '\n'
|
|
36
|
+
except AttributeError:
|
|
37
|
+
pass
|
|
38
|
+
try:
|
|
39
|
+
structured = structured + str(self.publisher) + ', '
|
|
40
|
+
except AttributeError:
|
|
41
|
+
pass
|
|
42
|
+
try:
|
|
43
|
+
structured = structured + str(self.doi) + ', '
|
|
44
|
+
except AttributeError:
|
|
45
|
+
pass
|
|
46
|
+
try:
|
|
47
|
+
structured = structured + str(self.year) +'.'
|
|
48
|
+
except AttributeError:
|
|
49
|
+
pass
|
|
50
|
+
return structured
|
|
51
|
+
|
|
52
|
+
@classmethod
|
|
53
|
+
def bibliography(cls, references: List['Reference'], structured: bool) -> str | List[Dict]:
|
|
54
|
+
"""Format a bibliography in a string from a list of references
|
|
55
|
+
|
|
56
|
+
:param List[Reference] references: List of References
|
|
57
|
+
:param bool structured: Wether we want a string bibliography or a list of references
|
|
58
|
+
|
|
59
|
+
:return: Bibliography in string or List format
|
|
60
|
+
"""
|
|
61
|
+
if structured:
|
|
62
|
+
return [vars(r) for r in references]
|
|
63
|
+
spacing = len(str(len(references)))
|
|
64
|
+
return '\n'.join([f"[{i+1:>{spacing}}] {str(reference)}\n" \
|
|
65
|
+
for i, reference in enumerate(references)])
|
iaml/shared_cache.py
ADDED
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
from collections import deque
|
|
2
|
+
from copy import deepcopy
|
|
3
|
+
from typing import Any, Tuple
|
|
4
|
+
import multiprocess.managers
|
|
5
|
+
from .logger import Logger
|
|
6
|
+
|
|
7
|
+
|
|
8
|
+
class CacheService:
|
|
9
|
+
"""Cache partagé, LRU bornée, clé=(fingerprint, df_hash)."""
|
|
10
|
+
def __init__(self, max_cache_size: int = 100) -> None:
|
|
11
|
+
self._saved = None
|
|
12
|
+
self._lru = None
|
|
13
|
+
self._max = max_cache_size
|
|
14
|
+
self._disabled = None
|
|
15
|
+
self._lock = None
|
|
16
|
+
|
|
17
|
+
def __set_backend__(self, saved, lru, disabled, lock, maxsize: int):
|
|
18
|
+
self._saved = saved
|
|
19
|
+
self._lru = lru
|
|
20
|
+
self._disabled = disabled
|
|
21
|
+
self._lock = lock
|
|
22
|
+
self._max = maxsize
|
|
23
|
+
|
|
24
|
+
# API
|
|
25
|
+
def disable(self) -> None:
|
|
26
|
+
with self._lock:
|
|
27
|
+
self._disabled.value = True
|
|
28
|
+
|
|
29
|
+
def enable(self) -> None:
|
|
30
|
+
with self._lock:
|
|
31
|
+
self._disabled.value = False
|
|
32
|
+
|
|
33
|
+
def get(self, fingerprint: str, df_hash: str) -> Any | None:
|
|
34
|
+
if self._disabled.value:
|
|
35
|
+
return None
|
|
36
|
+
key = (fingerprint, df_hash)
|
|
37
|
+
|
|
38
|
+
with self._lock:
|
|
39
|
+
if key in self._saved:
|
|
40
|
+
try:
|
|
41
|
+
self._lru.remove(key)
|
|
42
|
+
except ValueError:
|
|
43
|
+
pass
|
|
44
|
+
self._lru.append(key)
|
|
45
|
+
return deepcopy(self._saved[key])
|
|
46
|
+
return None
|
|
47
|
+
|
|
48
|
+
def put(self, fingerprint: str, df_hash: str, output: Any) -> None:
|
|
49
|
+
if self._disabled.value:
|
|
50
|
+
return
|
|
51
|
+
key = (fingerprint, df_hash)
|
|
52
|
+
|
|
53
|
+
with self._lock:
|
|
54
|
+
self._saved[key] = deepcopy(output)
|
|
55
|
+
try:
|
|
56
|
+
self._lru.remove(key)
|
|
57
|
+
except ValueError:
|
|
58
|
+
pass
|
|
59
|
+
self._lru.append(key)
|
|
60
|
+
# Éviction LRU
|
|
61
|
+
while len(self._lru) > self._max:
|
|
62
|
+
old_key = self._lru.pop(0)
|
|
63
|
+
self._saved.pop(old_key, None)
|
|
64
|
+
|
|
65
|
+
|
|
66
|
+
class CacheManager(multiprocess.managers.BaseManager):
|
|
67
|
+
pass
|
|
68
|
+
|
|
69
|
+
def start_cache_manager(max_cache_size: int = 100) -> tuple[CacheManager, CacheService]:
|
|
70
|
+
"""
|
|
71
|
+
Démarre un process manager et retourne (manager, cache_proxy).
|
|
72
|
+
À appeler UNE FOIS dans le process parent AVANT de lancer les workers.
|
|
73
|
+
"""
|
|
74
|
+
def _cache_factory():
|
|
75
|
+
from multiprocess.managers import SyncManager
|
|
76
|
+
sm = SyncManager()
|
|
77
|
+
sm.start()
|
|
78
|
+
saved = sm.dict()
|
|
79
|
+
lru = sm.list()
|
|
80
|
+
disabled = sm.Value('b', False)
|
|
81
|
+
lock = sm.RLock()
|
|
82
|
+
cache = CacheService(max_cache_size)
|
|
83
|
+
cache.__set_backend__(saved, lru, disabled, lock, max_cache_size)
|
|
84
|
+
return cache
|
|
85
|
+
|
|
86
|
+
CacheManager.register('Cache', callable=_cache_factory)
|
|
87
|
+
mgr = CacheManager()
|
|
88
|
+
mgr.start()
|
|
89
|
+
cache: CacheService = mgr.Cache()
|
|
90
|
+
return mgr, cache
|
|
@@ -0,0 +1,74 @@
|
|
|
1
|
+
"""A mandatory sklearn input transformer, fitted inside each candidate pipeline."""
|
|
2
|
+
from __future__ import annotations
|
|
3
|
+
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
from joblib import hash as joblib_hash
|
|
7
|
+
import numpy as np
|
|
8
|
+
import pandas as pd
|
|
9
|
+
from sklearn.base import clone
|
|
10
|
+
from sklearn.utils.validation import check_is_fitted
|
|
11
|
+
|
|
12
|
+
from .dataset import Dataset
|
|
13
|
+
from .step import Step
|
|
14
|
+
|
|
15
|
+
|
|
16
|
+
class SklearnPreprocessor(Step):
|
|
17
|
+
"""Adapt an unsupervised sklearn transformer to IAML's fit/transform protocol.
|
|
18
|
+
|
|
19
|
+
The caller's transformer is a template only. A fresh sklearn clone is fitted
|
|
20
|
+
on each supplied Dataset.X, without target values or patient-group columns.
|
|
21
|
+
``transformer_`` is the fitted clone, retained for prediction/provenance and
|
|
22
|
+
serialization. The template and its hyperparameters participate in cache
|
|
23
|
+
keys; the learned vocabulary does not alter the pipeline configuration.
|
|
24
|
+
|
|
25
|
+
A DataFrame output is required so column names and row alignment remain
|
|
26
|
+
explicit. This step is deliberately not registered under a search tag: it
|
|
27
|
+
can only enter the search through IAML's explicit initial_preprocessor.
|
|
28
|
+
"""
|
|
29
|
+
|
|
30
|
+
name = "Initial sklearn preprocessing"
|
|
31
|
+
can_be_disabled = False
|
|
32
|
+
|
|
33
|
+
def __init__(self, transformer: Any):
|
|
34
|
+
super().__init__()
|
|
35
|
+
if not callable(getattr(transformer, "fit", None)) or not callable(getattr(transformer, "transform", None)):
|
|
36
|
+
raise TypeError("initial_preprocessor must implement sklearn fit and transform")
|
|
37
|
+
self.transformer = clone(transformer)
|
|
38
|
+
self.configuration = {
|
|
39
|
+
"transformer_class": {"default": f"{type(transformer).__module__}.{type(transformer).__qualname__}"},
|
|
40
|
+
"transformer_parameters_hash": {"default": joblib_hash(self.transformer.get_params(deep=True))},
|
|
41
|
+
}
|
|
42
|
+
self.default_configuration()
|
|
43
|
+
self.is_interchangeable = False
|
|
44
|
+
self.optimizable = False
|
|
45
|
+
|
|
46
|
+
def fit(self, dataset: Dataset) -> "SklearnPreprocessor":
|
|
47
|
+
# Never retain categories learned by the generation sample or an earlier
|
|
48
|
+
# fold. sklearn.clone also removes fitted state supplied by the caller.
|
|
49
|
+
self.__dict__.pop("transformer_", None)
|
|
50
|
+
fitted = clone(self.transformer)
|
|
51
|
+
fitted.fit(dataset.X.copy(deep=True))
|
|
52
|
+
self.transformer_ = fitted
|
|
53
|
+
return self
|
|
54
|
+
|
|
55
|
+
def transform(self, X: pd.DataFrame) -> pd.DataFrame:
|
|
56
|
+
check_is_fitted(self, "transformer_")
|
|
57
|
+
if not isinstance(X, pd.DataFrame):
|
|
58
|
+
raise TypeError("IAML preprocessing requires a DataFrame input")
|
|
59
|
+
result = self.transformer_.transform(X.copy(deep=True))
|
|
60
|
+
if not isinstance(result, pd.DataFrame):
|
|
61
|
+
raise TypeError("initial_preprocessor must return a numeric DataFrame")
|
|
62
|
+
if len(result) != len(X) or not result.index.equals(X.index):
|
|
63
|
+
raise ValueError("initial_preprocessor must preserve row count, order and index")
|
|
64
|
+
if not result.columns.is_unique or not len(result.columns):
|
|
65
|
+
raise ValueError("initial_preprocessor must return unique, nonempty feature columns")
|
|
66
|
+
if not all(pd.api.types.is_numeric_dtype(dtype) for dtype in result.dtypes):
|
|
67
|
+
raise TypeError("initial_preprocessor output columns must all be numeric")
|
|
68
|
+
# IAML regards bool columns as categorical. Keep the explicit numeric
|
|
69
|
+
# representation, including BooleanDtype missing values, unambiguous.
|
|
70
|
+
boolean_columns = result.select_dtypes(include=["bool", "boolean"]).columns
|
|
71
|
+
if len(boolean_columns):
|
|
72
|
+
result = result.copy()
|
|
73
|
+
result[boolean_columns] = result[boolean_columns].astype(np.float32)
|
|
74
|
+
return result
|
|
@@ -0,0 +1,32 @@
|
|
|
1
|
+
"""Allow to split a dataset into n folds to compute crossvalidation"""
|
|
2
|
+
from typing import Iterator
|
|
3
|
+
from sklearn.model_selection import KFold as SKKFold, StratifiedKFold
|
|
4
|
+
from sklearn.model_selection import StratifiedGroupKFold
|
|
5
|
+
from sklearn.model_selection import GroupKFold
|
|
6
|
+
from ..dataset import Dataset
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
def kfold_splitter(dataset: Dataset, nb_folds: int = 5) -> Iterator[tuple['Dataset', 'Dataset']]:
|
|
10
|
+
"""Allow to split a dataset into n folds to compute crossvalidation
|
|
11
|
+
|
|
12
|
+
:param Dataset dataset: The dataset to split.
|
|
13
|
+
:param int, optional nb_folds: The number of folds to create. Default to 5.
|
|
14
|
+
:return: Iterator of tuples of train/test Dataset objects
|
|
15
|
+
"""
|
|
16
|
+
kwargs = {}
|
|
17
|
+
|
|
18
|
+
if dataset.type_of_target in ['binary', 'multiclass']:
|
|
19
|
+
if dataset.has_groups:
|
|
20
|
+
kfold = StratifiedGroupKFold(nb_folds)
|
|
21
|
+
kwargs['groups'] = dataset.groups
|
|
22
|
+
else:
|
|
23
|
+
kfold = StratifiedKFold(nb_folds)
|
|
24
|
+
else:
|
|
25
|
+
if dataset.has_groups:
|
|
26
|
+
kfold = GroupKFold(nb_folds)
|
|
27
|
+
kwargs['groups'] = dataset.groups
|
|
28
|
+
else:
|
|
29
|
+
kfold = SKKFold(nb_folds)
|
|
30
|
+
|
|
31
|
+
for ds_train, ds_test in dataset.split(kfold.split, **kwargs):
|
|
32
|
+
yield (ds_train, ds_test)
|
|
@@ -0,0 +1,26 @@
|
|
|
1
|
+
"""Allow to split randomly a dataset to train/test"""
|
|
2
|
+
from typing import Iterator
|
|
3
|
+
from sklearn.model_selection import ShuffleSplit, GroupShuffleSplit
|
|
4
|
+
from ..dataset import Dataset
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def random_splitter(
|
|
8
|
+
dataset: Dataset,
|
|
9
|
+
ratio: float = 0.2,
|
|
10
|
+
random_state: int = 42) -> Iterator[tuple['Dataset', 'Dataset']]:
|
|
11
|
+
"""Allow to split randomly a dataset to train/test
|
|
12
|
+
|
|
13
|
+
:param Dataset dataset: The dataset to split.
|
|
14
|
+
:param float, optional ratio: The train/test ratio. Default to 0.2.
|
|
15
|
+
:param int, optional random_state: The random seed used. Default to 42.
|
|
16
|
+
:return: Iterator of tuples of train/test Dataset objects
|
|
17
|
+
"""
|
|
18
|
+
kwargs = {}
|
|
19
|
+
if dataset.has_groups:
|
|
20
|
+
kwargs['groups'] = dataset.groups
|
|
21
|
+
splitter = GroupShuffleSplit(1, test_size=ratio, random_state=random_state)
|
|
22
|
+
else:
|
|
23
|
+
splitter = ShuffleSplit(1, test_size=ratio, random_state=random_state)
|
|
24
|
+
|
|
25
|
+
for ds_train, ds_test in dataset.split(splitter.split, **kwargs):
|
|
26
|
+
yield (ds_train, ds_test)
|
iaml/stack.py
ADDED
|
@@ -0,0 +1,39 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Frozen version of Step created to be stacked in an Candidate
|
|
3
|
+
"""
|
|
4
|
+
|
|
5
|
+
from typing import Dict
|
|
6
|
+
from .step import Step
|
|
7
|
+
class Stack:
|
|
8
|
+
"""
|
|
9
|
+
Frozen version of Step created to be stacked in an Candidate
|
|
10
|
+
|
|
11
|
+
:param Step step_class: The step class used to create this stack
|
|
12
|
+
:param Dict configuration: The configuration dictionnary
|
|
13
|
+
:param int step_id: The id of the Step
|
|
14
|
+
"""
|
|
15
|
+
def __init__(self, step_class: Step, configuration: Dict, step_id: int) -> None:
|
|
16
|
+
"""Initialize a stack
|
|
17
|
+
"""
|
|
18
|
+
self.step_class: Step = step_class
|
|
19
|
+
"""The step used in this Stack"""
|
|
20
|
+
|
|
21
|
+
self.configuration: dict = configuration
|
|
22
|
+
"""The step configuration"""
|
|
23
|
+
|
|
24
|
+
self.step_id: int = step_id
|
|
25
|
+
"""The step id"""
|
|
26
|
+
|
|
27
|
+
def __str__(self) -> str:
|
|
28
|
+
"""Return the step name as Stack representation
|
|
29
|
+
|
|
30
|
+
:return: The step name
|
|
31
|
+
"""
|
|
32
|
+
return self.step_class.name
|
|
33
|
+
|
|
34
|
+
def explain(self) -> str:
|
|
35
|
+
"""Return explanation string of Step
|
|
36
|
+
|
|
37
|
+
:return: The explanation of the Step in string format
|
|
38
|
+
"""
|
|
39
|
+
return self.step_class.explain(self.configuration)
|