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,91 @@
|
|
|
1
|
+
"""[STEP] Random Survival Forest"""
|
|
2
|
+
import textwrap
|
|
3
|
+
from typing import Any
|
|
4
|
+
from sksurv.ensemble import RandomSurvivalForest
|
|
5
|
+
|
|
6
|
+
from ....predictor import Predictor
|
|
7
|
+
from ....candidate import Candidate
|
|
8
|
+
from ....dataset import Dataset
|
|
9
|
+
from ....decorators.all import is_step
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@is_step('predictor', 'tabular', 'survival')
|
|
13
|
+
class ActRandomSurvivalForest(Predictor):
|
|
14
|
+
"""[STEP] Random Survival Forest"""
|
|
15
|
+
|
|
16
|
+
name: str = "RandomSurvivalForest"
|
|
17
|
+
_usage: str = "Use when you want a flexible tree-ensemble survival model; choose over ActCox when PH is doubtful, or consider ActExtraSurvivalTrees for more randomness. Applicable to tabular time-to-event data with censoring. Avoid when data are tiny or effects are well modeled by linear PH."
|
|
18
|
+
_description: str = textwrap.dedent('''\
|
|
19
|
+
RandomSurvivalForest is a survival analysis algorithm
|
|
20
|
+
that uses an ensemble of decision trees to estimate the survival function
|
|
21
|
+
over time. It is a non-parametric model that handles complex relationships
|
|
22
|
+
and can model non-linear effects of covariates.''')
|
|
23
|
+
_description_long: str = textwrap.dedent('''\
|
|
24
|
+
RandomSurvivalForest is a flexible survival analysis
|
|
25
|
+
algorithm that uses an ensemble of decision trees to predict the time
|
|
26
|
+
until an event occurs. It extends the concept of random forests to survival
|
|
27
|
+
data, handling complex interactions and non-linear relationships between
|
|
28
|
+
input features (covariates). Unlike parametric models such as the Cox
|
|
29
|
+
Proportional Hazards model, RandomSurvivalForest makes fewer assumptions
|
|
30
|
+
about the underlying data, making it useful in cases where the assumptions
|
|
31
|
+
of proportional hazards do not hold. It also efficiently manages censored
|
|
32
|
+
data, where the event may not have occurred during the study period.''')
|
|
33
|
+
refs: list[dict[str, Any]] = [
|
|
34
|
+
{
|
|
35
|
+
'year': 2008,
|
|
36
|
+
'name': 'Random survival forests',
|
|
37
|
+
'authors': [
|
|
38
|
+
'H. Ishwaran',
|
|
39
|
+
'U. B. Kogalur',
|
|
40
|
+
'E. H. Blackstone',
|
|
41
|
+
'M. S. Lauer'
|
|
42
|
+
],
|
|
43
|
+
'doi': 'https://doi.org/10.1214/08-AOAS169',
|
|
44
|
+
'publisher': 'The Annals of Applied Statistics, 2(3): page 841-860'
|
|
45
|
+
}
|
|
46
|
+
]
|
|
47
|
+
|
|
48
|
+
def __init__(self):
|
|
49
|
+
self.configuration: dict = {
|
|
50
|
+
'n_estimators': {
|
|
51
|
+
'description': 'Number of trees in the forest.',
|
|
52
|
+
'default': 100,
|
|
53
|
+
'range': [1, 1000],
|
|
54
|
+
'passthrough': True
|
|
55
|
+
},
|
|
56
|
+
'min_samples_split': {
|
|
57
|
+
'description': 'The minimum number of samples required to split an internal node.',
|
|
58
|
+
'default': 2,
|
|
59
|
+
'range': [2, 20],
|
|
60
|
+
'passthrough': True
|
|
61
|
+
},
|
|
62
|
+
'min_samples_leaf': {
|
|
63
|
+
'description': 'The minimum number of samples required to be at a leaf node.',
|
|
64
|
+
'default': 1,
|
|
65
|
+
'range': [1, 20],
|
|
66
|
+
'passthrough': True
|
|
67
|
+
},
|
|
68
|
+
'max_depth': {
|
|
69
|
+
'description': textwrap.dedent('''\
|
|
70
|
+
The maximum depth of the tree. If None, then nodes are
|
|
71
|
+
expanded until all leaves are pure.'''),
|
|
72
|
+
'default': None,
|
|
73
|
+
'range': [1, None],
|
|
74
|
+
'passthrough': True
|
|
75
|
+
}
|
|
76
|
+
}
|
|
77
|
+
self.model: RandomSurvivalForest = None
|
|
78
|
+
|
|
79
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
80
|
+
self.model = RandomSurvivalForest(
|
|
81
|
+
**self.passthrough_parameters()
|
|
82
|
+
)
|
|
83
|
+
X, y = dataset.to_survival()
|
|
84
|
+
self.model.fit(X, y)
|
|
85
|
+
return self
|
|
86
|
+
|
|
87
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
88
|
+
return dataset.type_of_target == 'survival'
|
|
89
|
+
|
|
90
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
91
|
+
return 0.5 # neutral
|
|
@@ -0,0 +1,80 @@
|
|
|
1
|
+
"""[STEP] Componentwise Gradient Boosting Survival Analysis"""
|
|
2
|
+
import textwrap
|
|
3
|
+
from typing import Any
|
|
4
|
+
from sksurv.ensemble import ComponentwiseGradientBoostingSurvivalAnalysis
|
|
5
|
+
|
|
6
|
+
from ....predictor import Predictor
|
|
7
|
+
from ....candidate import Candidate
|
|
8
|
+
from ....dataset import Dataset
|
|
9
|
+
from ....decorators.all import is_step
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@is_step('predictor', 'tabular', 'survival', 'minimal_predictor')
|
|
13
|
+
class ActComponentwiseGradientBoostingSurvivalAnalysis(Predictor):
|
|
14
|
+
"""[STEP] Componentwise Gradient Boosting Survival Analysis"""
|
|
15
|
+
|
|
16
|
+
name: str = "ComponentwiseGradientBoostingSurvivalAnalysis"
|
|
17
|
+
_usage: str = "Use when you want stagewise boosting with feature selection for survival instead of ActGradientBoostingSurvivalAnalysis. Applicable to tabular censored survival data, especially with many features. Avoid when a linear model like ActCox or ActCoxnetSurvivalAnalysis is preferred."
|
|
18
|
+
_description: str = textwrap.dedent('''\
|
|
19
|
+
ComponentwiseGradientBoostingSurvivalAnalysis is a survival analysis
|
|
20
|
+
algorithm that uses gradient boosting with componentwise (stagewise) updates
|
|
21
|
+
to estimate the survival function over time. This variant of boosting allows
|
|
22
|
+
the model to fit individual components (features) in a stagewise manner,
|
|
23
|
+
making it a more interpretable approach for feature selection and model
|
|
24
|
+
refinement in survival analysis.''')
|
|
25
|
+
_description_long: str = textwrap.dedent('''\
|
|
26
|
+
ComponentwiseGradientBoostingSurvivalAnalysis extends
|
|
27
|
+
gradient boosting for survival analysis by applying updates one component
|
|
28
|
+
(feature) at a time. This approach improves the model's ability to handle
|
|
29
|
+
sparse datasets or datasets with high-dimensional features, where only
|
|
30
|
+
a few variables may have significant effects on survival outcomes.
|
|
31
|
+
It provides a more interpretable framework for survival analysis,
|
|
32
|
+
as each boosting iteration focuses on fitting individual covariates
|
|
33
|
+
rather than combining all features at once. This method is particularly
|
|
34
|
+
suited for feature selection and handling censored survival data.''')
|
|
35
|
+
refs: list[dict[str, Any]] = [
|
|
36
|
+
{
|
|
37
|
+
'year': 2006,
|
|
38
|
+
'name': 'Survival ensembles',
|
|
39
|
+
'authors': [
|
|
40
|
+
'T. Hothorn',
|
|
41
|
+
'P. B5hlmann',
|
|
42
|
+
'S. Dudoit',
|
|
43
|
+
'A. Molinaro',
|
|
44
|
+
'M. J. van der Laan'
|
|
45
|
+
],
|
|
46
|
+
'doi': 'https://doi.org/10.1093/biostatistics/kxj011',
|
|
47
|
+
'publisher': 'Biostatistics, 7(3), 355-373'
|
|
48
|
+
}
|
|
49
|
+
]
|
|
50
|
+
|
|
51
|
+
def __init__(self):
|
|
52
|
+
self.configuration: dict = {
|
|
53
|
+
'n_estimators': {
|
|
54
|
+
'description': 'Number of boosting stages to be run.',
|
|
55
|
+
'default': 100,
|
|
56
|
+
'range': [1, 1000],
|
|
57
|
+
'passthrough': True
|
|
58
|
+
},
|
|
59
|
+
'learning_rate': {
|
|
60
|
+
'description': 'Learning rate shrinks the contribution of each component.',
|
|
61
|
+
'default': 0.1,
|
|
62
|
+
'range': [0.01, 1.0],
|
|
63
|
+
'passthrough': True
|
|
64
|
+
}
|
|
65
|
+
}
|
|
66
|
+
self.model: ComponentwiseGradientBoostingSurvivalAnalysis = None
|
|
67
|
+
|
|
68
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
69
|
+
self.model = ComponentwiseGradientBoostingSurvivalAnalysis(
|
|
70
|
+
**self.passthrough_parameters()
|
|
71
|
+
)
|
|
72
|
+
X, y = dataset.to_survival()
|
|
73
|
+
self.model.fit(X, y)
|
|
74
|
+
return self
|
|
75
|
+
|
|
76
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
77
|
+
return dataset.type_of_target == 'survival'
|
|
78
|
+
|
|
79
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
80
|
+
return 0.5 # neutral
|
|
@@ -0,0 +1,120 @@
|
|
|
1
|
+
"""[STEP] SurvivalTree"""
|
|
2
|
+
import textwrap
|
|
3
|
+
from typing import Any
|
|
4
|
+
from sksurv.tree import SurvivalTree
|
|
5
|
+
|
|
6
|
+
from ....predictor import Predictor
|
|
7
|
+
from ....candidate import Candidate
|
|
8
|
+
from ....dataset import Dataset
|
|
9
|
+
from ....decorators.all import is_step
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
@is_step('predictor', 'tabular', 'survival', 'baseline_predictor')
|
|
13
|
+
class ActSurvivalTree(Predictor):
|
|
14
|
+
"""[STEP] SurvivalTree"""
|
|
15
|
+
|
|
16
|
+
name: str = "SurvivalTree"
|
|
17
|
+
_description: str = textwrap.dedent('''\
|
|
18
|
+
SurvivalTree is a decision tree algorithm tailored for survival analysis.
|
|
19
|
+
It constructs a tree structure based on the log-rank test, where each split
|
|
20
|
+
is designed to separate data by survival times. The method handles censored data
|
|
21
|
+
and outputs risk scores, cumulative hazard functions, and survival functions
|
|
22
|
+
based on the tree's terminal nodes.''')
|
|
23
|
+
_description_long: str = textwrap.dedent('''\
|
|
24
|
+
SurvivalTree builds a decision tree based on survival data, using the
|
|
25
|
+
log-rank splitting rule to determine the best splits. It is a non-parametric model that
|
|
26
|
+
is particularly suited for survival analysis with right-censored data. The model provides
|
|
27
|
+
both cumulative hazard and survival functions at each terminal node, making it useful for
|
|
28
|
+
clinical risk prediction and other applications where time-to-event outcomes are crucial.
|
|
29
|
+
''')
|
|
30
|
+
_usage: str = "Use when you need an interpretable survival baseline; compare ActCox or ActExtraSurvivalTrees for linear or ensemble options. Applicable to tabular time-to-event data with right censoring. Avoid when higher accuracy is required or proportional-hazards structure is assumed."
|
|
31
|
+
refs: list[dict[str, Any]] = [
|
|
32
|
+
{
|
|
33
|
+
'year': 1993,
|
|
34
|
+
'name': 'Survival Trees by Goodness of Split',
|
|
35
|
+
'authors': [
|
|
36
|
+
'M. Leblanc',
|
|
37
|
+
'J. Crowley'
|
|
38
|
+
],
|
|
39
|
+
'doi': 'https://doi.org/10.1080/01621459.1993.10476296',
|
|
40
|
+
'publisher': 'Journal of the American Statistical Association, 88(422), 457-467'
|
|
41
|
+
}
|
|
42
|
+
]
|
|
43
|
+
|
|
44
|
+
def __init__(self):
|
|
45
|
+
self.configuration: dict = {
|
|
46
|
+
'splitter': {
|
|
47
|
+
'description': textwrap.dedent('''\
|
|
48
|
+
The strategy used to split at each node. Supported: "best",
|
|
49
|
+
"random".'''),
|
|
50
|
+
'default': 'best',
|
|
51
|
+
'categorical': ['best', 'random'],
|
|
52
|
+
'passthrough': True
|
|
53
|
+
},
|
|
54
|
+
'max_depth': {
|
|
55
|
+
'description': 'The maximum depth of the tree.',
|
|
56
|
+
'default': None,
|
|
57
|
+
'range': [1, None],
|
|
58
|
+
'passthrough': True
|
|
59
|
+
},
|
|
60
|
+
'min_samples_split': {
|
|
61
|
+
'description': 'The minimum number of samples required to split an internal node.',
|
|
62
|
+
'default': 6,
|
|
63
|
+
'range': [2, 20],
|
|
64
|
+
'passthrough': True
|
|
65
|
+
},
|
|
66
|
+
'min_samples_leaf': {
|
|
67
|
+
'description': 'The minimum number of samples required to be at a leaf node.',
|
|
68
|
+
'default': 3,
|
|
69
|
+
'range': [1, 20],
|
|
70
|
+
'passthrough': True
|
|
71
|
+
},
|
|
72
|
+
'min_weight_fraction_leaf': {
|
|
73
|
+
'description': textwrap.dedent('''\
|
|
74
|
+
The minimum weighted fraction of the input samples required
|
|
75
|
+
to be at a leaf node.'''),
|
|
76
|
+
'default': 0.0,
|
|
77
|
+
'range': [0.0, 0.5],
|
|
78
|
+
'passthrough': True
|
|
79
|
+
},
|
|
80
|
+
'max_features': {
|
|
81
|
+
'description': textwrap.dedent('''\
|
|
82
|
+
The number of features to consider when looking for the
|
|
83
|
+
best split.'''),
|
|
84
|
+
'default': None,
|
|
85
|
+
'categorical': [None, 'sqrt', 'log2'],
|
|
86
|
+
'passthrough': True
|
|
87
|
+
},
|
|
88
|
+
'random_state': {
|
|
89
|
+
'description': 'Random seed (integer or None) for the estimator.',
|
|
90
|
+
'default': None,
|
|
91
|
+
'passthrough': True
|
|
92
|
+
},
|
|
93
|
+
'max_leaf_nodes': {
|
|
94
|
+
'description': 'Grow a tree with a maximum number of leaf nodes.',
|
|
95
|
+
'default': None,
|
|
96
|
+
'range': [None, 1000],
|
|
97
|
+
'passthrough': True
|
|
98
|
+
},
|
|
99
|
+
'low_memory': {
|
|
100
|
+
'description': 'Reduce memory usage but disable some prediction functions.',
|
|
101
|
+
'default': False,
|
|
102
|
+
'categorical': [True, False],
|
|
103
|
+
'passthrough': True
|
|
104
|
+
}
|
|
105
|
+
}
|
|
106
|
+
self.model: SurvivalTree = None
|
|
107
|
+
|
|
108
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
109
|
+
self.model = SurvivalTree(
|
|
110
|
+
**self.passthrough_parameters()
|
|
111
|
+
)
|
|
112
|
+
X, y = dataset.to_survival()
|
|
113
|
+
self.model.fit(X, y)
|
|
114
|
+
return self
|
|
115
|
+
|
|
116
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
117
|
+
return dataset.type_of_target == 'survival'
|
|
118
|
+
|
|
119
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
120
|
+
return 0.5 # neutral
|
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
"""Compatibility import for the former survival boosting module.
|
|
2
|
+
|
|
3
|
+
The implementation lives in ``act_gradient_boosting_survival_analysis``.
|
|
4
|
+
Re-export the same class so existing imports and pickles keep resolving without
|
|
5
|
+
registering a second predictor. The backend is scikit-survival, not XGBoost.
|
|
6
|
+
"""
|
|
7
|
+
from .act_gradient_boosting_survival_analysis import ActGradientBoostingSurvivalAnalysis
|
|
8
|
+
|
|
9
|
+
__all__ = ['ActGradientBoostingSurvivalAnalysis']
|
|
@@ -0,0 +1,230 @@
|
|
|
1
|
+
"""[STEP] Weibull AFT"""
|
|
2
|
+
import inspect
|
|
3
|
+
import textwrap
|
|
4
|
+
from typing import Any
|
|
5
|
+
|
|
6
|
+
import numpy as np
|
|
7
|
+
|
|
8
|
+
try:
|
|
9
|
+
from sksurv.linear_model import WeibullAFT
|
|
10
|
+
except Exception: # pragma: no cover - optional dependency
|
|
11
|
+
try:
|
|
12
|
+
from sksurv.parametric import WeibullAFT
|
|
13
|
+
except Exception: # pragma: no cover - optional dependency
|
|
14
|
+
WeibullAFT = None
|
|
15
|
+
|
|
16
|
+
try:
|
|
17
|
+
from lifelines import WeibullAFTFitter
|
|
18
|
+
except Exception: # pragma: no cover - optional dependency
|
|
19
|
+
WeibullAFTFitter = None
|
|
20
|
+
|
|
21
|
+
from ....predictor import Predictor
|
|
22
|
+
from ....candidate import Candidate
|
|
23
|
+
from ....dataset import Dataset
|
|
24
|
+
from ....data_type import DataType
|
|
25
|
+
from ....decorators.all import is_step
|
|
26
|
+
|
|
27
|
+
|
|
28
|
+
@is_step('predictor', 'tabular', 'survival')
|
|
29
|
+
class ActWeibullAFT(Predictor):
|
|
30
|
+
"""[STEP] Weibull AFT"""
|
|
31
|
+
|
|
32
|
+
name: str = "WeibullAFT"
|
|
33
|
+
_usage: str = "Use when you want a parametric Weibull AFT model with time ratios, rather than ActCox. Applicable to tabular right-censored survival data with numeric features. Avoid when Weibull fit is implausible or you need flexible nonparametric models like ActExtraSurvivalTrees."
|
|
34
|
+
_description: str = textwrap.dedent('''\
|
|
35
|
+
WeibullAFT is a parametric accelerated failure time model
|
|
36
|
+
that assumes survival times follow a Weibull distribution and
|
|
37
|
+
models how covariates speed up or slow down the event time.
|
|
38
|
+
Uses scikit-survival when available and falls back to lifelines.''')
|
|
39
|
+
_description_long: str = textwrap.dedent('''\
|
|
40
|
+
WeibullAFT fits an accelerated failure time model where the log of
|
|
41
|
+
survival time is a linear function of the input features and the
|
|
42
|
+
baseline survival follows a Weibull distribution. The model provides
|
|
43
|
+
parametric survival and hazard estimates, supports right-censored data,
|
|
44
|
+
and yields interpretable covariate effects on time-to-event outcomes.
|
|
45
|
+
This step uses scikit-survival when available and falls back to
|
|
46
|
+
lifelines when needed.''')
|
|
47
|
+
|
|
48
|
+
def __init__(self):
|
|
49
|
+
self.configuration: dict = {
|
|
50
|
+
'alpha': {
|
|
51
|
+
'description': textwrap.dedent('''\
|
|
52
|
+
Regularization strength for scikit-survival. Higher values
|
|
53
|
+
enforce stronger regularization.'''),
|
|
54
|
+
'default': 0.05,
|
|
55
|
+
'range': [1e-04, 10.0]
|
|
56
|
+
},
|
|
57
|
+
'penalizer': {
|
|
58
|
+
'description': textwrap.dedent('''\
|
|
59
|
+
L2 penalizer strength for lifelines models.'''),
|
|
60
|
+
'default': 0.0,
|
|
61
|
+
'range': [0.0, 10.0]
|
|
62
|
+
},
|
|
63
|
+
'l1_ratio': {
|
|
64
|
+
'description': textwrap.dedent('''\
|
|
65
|
+
Mixing parameter between L1 and L2 penalties.'''),
|
|
66
|
+
'default': 0.0,
|
|
67
|
+
'range': [0.0, 1.0]
|
|
68
|
+
},
|
|
69
|
+
'fit_intercept': {
|
|
70
|
+
'description': 'Whether to fit the intercept term.',
|
|
71
|
+
'default': True,
|
|
72
|
+
'categorical': [True, False]
|
|
73
|
+
},
|
|
74
|
+
'max_iter': {
|
|
75
|
+
'description': 'Maximum number of iterations for the optimizer.',
|
|
76
|
+
'default': 1000,
|
|
77
|
+
'range': [10, 100000]
|
|
78
|
+
},
|
|
79
|
+
'tol': {
|
|
80
|
+
'description': 'Stopping tolerance.',
|
|
81
|
+
'default': 1e-07,
|
|
82
|
+
'range': [1e-09, 1e-03]
|
|
83
|
+
}
|
|
84
|
+
}
|
|
85
|
+
self.model: Any = None
|
|
86
|
+
self.columns: list[str] = []
|
|
87
|
+
self.backend: str | None = None
|
|
88
|
+
self._lifelines_event_col: str | None = None
|
|
89
|
+
self._lifelines_time_col: str | None = None
|
|
90
|
+
|
|
91
|
+
def _select_features(self, X):
|
|
92
|
+
if self.columns and hasattr(X, 'columns'):
|
|
93
|
+
return X[self.columns]
|
|
94
|
+
return X
|
|
95
|
+
|
|
96
|
+
def _resolve_backend(self) -> tuple[str | None, Any | None]:
|
|
97
|
+
if WeibullAFT is not None:
|
|
98
|
+
return 'sksurv', WeibullAFT
|
|
99
|
+
if WeibullAFTFitter is not None:
|
|
100
|
+
return 'lifelines', WeibullAFTFitter
|
|
101
|
+
return None, None
|
|
102
|
+
|
|
103
|
+
def _model_parameters(self, model_cls: Any, backend: str) -> dict[str, Any]:
|
|
104
|
+
params = self.passthrough_parameters()
|
|
105
|
+
if backend == 'lifelines':
|
|
106
|
+
params.pop('alpha', None)
|
|
107
|
+
|
|
108
|
+
if model_cls is None:
|
|
109
|
+
return params
|
|
110
|
+
|
|
111
|
+
try:
|
|
112
|
+
sig_params = inspect.signature(model_cls).parameters
|
|
113
|
+
except (TypeError, ValueError):
|
|
114
|
+
return params
|
|
115
|
+
|
|
116
|
+
return {key: value for key, value in params.items() if key in sig_params}
|
|
117
|
+
|
|
118
|
+
@staticmethod
|
|
119
|
+
def _split_survival_target(y) -> tuple[np.ndarray, np.ndarray]:
|
|
120
|
+
samples = Dataset.normalize_survival_target(y)
|
|
121
|
+
if not samples:
|
|
122
|
+
return np.array([], dtype=bool), np.array([], dtype=float)
|
|
123
|
+
events, times = zip(*samples)
|
|
124
|
+
return np.asarray(events, dtype=bool), np.asarray(times, dtype=float)
|
|
125
|
+
|
|
126
|
+
@staticmethod
|
|
127
|
+
def _unique_column_name(base: str, columns) -> str:
|
|
128
|
+
name = base
|
|
129
|
+
while name in columns:
|
|
130
|
+
name = f"_{name}"
|
|
131
|
+
return name
|
|
132
|
+
|
|
133
|
+
def _prepare_lifelines_frame(self, X, y) -> tuple[Any, str, str]:
|
|
134
|
+
X_selected = self._select_features(X).copy()
|
|
135
|
+
event_col = self._unique_column_name("_event", X_selected.columns)
|
|
136
|
+
time_col = self._unique_column_name("_time", X_selected.columns)
|
|
137
|
+
events, times = self._split_survival_target(y)
|
|
138
|
+
X_selected[event_col] = events
|
|
139
|
+
X_selected[time_col] = times
|
|
140
|
+
return X_selected, time_col, event_col
|
|
141
|
+
|
|
142
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
143
|
+
self.columns = dataset.get_columns_names_by_type(DataType.NUMERIC)
|
|
144
|
+
if not self.columns:
|
|
145
|
+
self.columns = dataset.features
|
|
146
|
+
|
|
147
|
+
backend, model_cls = self._resolve_backend()
|
|
148
|
+
if model_cls is None:
|
|
149
|
+
raise ImportError(
|
|
150
|
+
"WeibullAFT requires scikit-survival or lifelines to be installed."
|
|
151
|
+
)
|
|
152
|
+
|
|
153
|
+
self.backend = backend
|
|
154
|
+
self.model = model_cls(**self._model_parameters(model_cls, backend))
|
|
155
|
+
|
|
156
|
+
if backend == 'sksurv':
|
|
157
|
+
X, y = dataset.to_survival()
|
|
158
|
+
self.model.fit(self._select_features(X), y)
|
|
159
|
+
else:
|
|
160
|
+
X, time_col, event_col = self._prepare_lifelines_frame(dataset.X, dataset.y)
|
|
161
|
+
self.model.fit(X, duration_col=time_col, event_col=event_col)
|
|
162
|
+
self._lifelines_time_col = time_col
|
|
163
|
+
self._lifelines_event_col = event_col
|
|
164
|
+
|
|
165
|
+
return self
|
|
166
|
+
|
|
167
|
+
def predict(self, X):
|
|
168
|
+
X_selected = self._select_features(X)
|
|
169
|
+
if self.model and hasattr(self.model, 'predict'):
|
|
170
|
+
return super().predict(X_selected)
|
|
171
|
+
if self.model and hasattr(self.model, 'predict_median'):
|
|
172
|
+
return self.model.predict_median(X_selected)
|
|
173
|
+
if self.model and hasattr(self.model, 'predict_expectation'):
|
|
174
|
+
return self.model.predict_expectation(X_selected)
|
|
175
|
+
return None
|
|
176
|
+
|
|
177
|
+
def predict_survival_function(self, X):
|
|
178
|
+
X_selected = self._select_features(X)
|
|
179
|
+
if self.model and hasattr(self.model, 'predict_survival_function'):
|
|
180
|
+
return self.model.predict_survival_function(X_selected)
|
|
181
|
+
raise AttributeError("Unable to predict survival function with this model")
|
|
182
|
+
|
|
183
|
+
def predict_cumulative_hazard_function(self, X):
|
|
184
|
+
X_selected = self._select_features(X)
|
|
185
|
+
if self.model and hasattr(self.model, 'predict_cumulative_hazard_function'):
|
|
186
|
+
return self.model.predict_cumulative_hazard_function(X_selected)
|
|
187
|
+
if self.model and hasattr(self.model, 'predict_cumulative_hazard'):
|
|
188
|
+
return self.model.predict_cumulative_hazard(X_selected)
|
|
189
|
+
raise AttributeError("Unable to predict cumulative hazard function with this model")
|
|
190
|
+
|
|
191
|
+
def score(self, X, y=None, *args, **kwargs):
|
|
192
|
+
if self.backend == 'lifelines':
|
|
193
|
+
if y is None:
|
|
194
|
+
if (
|
|
195
|
+
self._lifelines_event_col is None
|
|
196
|
+
or self._lifelines_time_col is None
|
|
197
|
+
or not hasattr(X, 'columns')
|
|
198
|
+
or self._lifelines_event_col not in X.columns
|
|
199
|
+
or self._lifelines_time_col not in X.columns
|
|
200
|
+
):
|
|
201
|
+
raise ValueError(
|
|
202
|
+
"lifelines score requires y or a DataFrame containing the "
|
|
203
|
+
"duration/event columns from fit."
|
|
204
|
+
)
|
|
205
|
+
X_frame = self._select_features(X).copy()
|
|
206
|
+
X_frame[self._lifelines_event_col] = X[self._lifelines_event_col]
|
|
207
|
+
X_frame[self._lifelines_time_col] = X[self._lifelines_time_col]
|
|
208
|
+
return self.model.score(X_frame, *args, **kwargs)
|
|
209
|
+
X_frame = self._select_features(X).copy()
|
|
210
|
+
events, times = self._split_survival_target(y)
|
|
211
|
+
if self._lifelines_event_col is None or self._lifelines_time_col is None:
|
|
212
|
+
self._lifelines_event_col = self._unique_column_name(
|
|
213
|
+
"_event", X_frame.columns
|
|
214
|
+
)
|
|
215
|
+
self._lifelines_time_col = self._unique_column_name(
|
|
216
|
+
"_time", X_frame.columns
|
|
217
|
+
)
|
|
218
|
+
X_frame[self._lifelines_event_col] = events
|
|
219
|
+
X_frame[self._lifelines_time_col] = times
|
|
220
|
+
return self.model.score(X_frame, *args, **kwargs)
|
|
221
|
+
return self.model.score(self._select_features(X), y, *args, **kwargs)
|
|
222
|
+
|
|
223
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
224
|
+
backend_available = WeibullAFT is not None or WeibullAFTFitter is not None
|
|
225
|
+
return dataset.type_of_target == 'survival' \
|
|
226
|
+
and bool(dataset.get_columns_names_by_type(DataType.NUMERIC)) \
|
|
227
|
+
and backend_available
|
|
228
|
+
|
|
229
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
230
|
+
return 0.5 # neutral
|
iaml/cache.py
ADDED
|
@@ -0,0 +1,61 @@
|
|
|
1
|
+
# cache_singleton.py
|
|
2
|
+
from typing import Any
|
|
3
|
+
import pickle
|
|
4
|
+
import pandas as pd
|
|
5
|
+
|
|
6
|
+
from .meta_singleton import MetaSingleton
|
|
7
|
+
from .cache_keys import hash_df
|
|
8
|
+
from .dataset import Dataset
|
|
9
|
+
from .logger import Logger
|
|
10
|
+
|
|
11
|
+
|
|
12
|
+
def _data_fingerprint(dataset: Dataset | pd.DataFrame | str) -> str:
|
|
13
|
+
"""Accept complete datasets, frozen keys, and legacy feature-only inputs."""
|
|
14
|
+
if isinstance(dataset, str):
|
|
15
|
+
return dataset
|
|
16
|
+
if isinstance(dataset, Dataset):
|
|
17
|
+
return dataset.fingerprint()
|
|
18
|
+
return hash_df(dataset)
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class Cache(metaclass=MetaSingleton):
|
|
22
|
+
"""Façade singleton vers un cache partagé."""
|
|
23
|
+
def __init__(self) -> None:
|
|
24
|
+
self._backend = None
|
|
25
|
+
|
|
26
|
+
def configure(self, backend) -> None:
|
|
27
|
+
"""Brancher le proxy partagé (à faire dans le parent ET dans chaque worker)."""
|
|
28
|
+
self._backend = backend
|
|
29
|
+
|
|
30
|
+
def disable(self) -> None:
|
|
31
|
+
if self._backend: self._backend.disable()
|
|
32
|
+
|
|
33
|
+
def enable(self) -> None:
|
|
34
|
+
if self._backend: self._backend.enable()
|
|
35
|
+
|
|
36
|
+
def from_cache(self, fingerprint: str, dataset: Dataset | pd.DataFrame | str) -> Any | None:
|
|
37
|
+
"""Look up an operation using all dataset inputs or a precomputed key.
|
|
38
|
+
|
|
39
|
+
DataFrame inputs remain supported for feature-only operations. Supervised
|
|
40
|
+
operations must pass a Dataset, or its fingerprint captured before mutation.
|
|
41
|
+
"""
|
|
42
|
+
if not self._backend:
|
|
43
|
+
return None
|
|
44
|
+
try:
|
|
45
|
+
return self._backend.get(fingerprint, _data_fingerprint(dataset))
|
|
46
|
+
except (EOFError, OSError) as exc:
|
|
47
|
+
Logger().warning(f"Shared cache disabled after get failure: {exc!r}")
|
|
48
|
+
self._backend = None
|
|
49
|
+
return None
|
|
50
|
+
|
|
51
|
+
def add_to_cache(
|
|
52
|
+
self, fingerprint: str, dataset: Dataset | pd.DataFrame | str, output: Any
|
|
53
|
+
) -> None:
|
|
54
|
+
"""Store an operation result under the same input key used for lookup."""
|
|
55
|
+
if not self._backend:
|
|
56
|
+
return
|
|
57
|
+
try:
|
|
58
|
+
self._backend.put(fingerprint, _data_fingerprint(dataset), output)
|
|
59
|
+
except (EOFError, OSError, pickle.PicklingError, TypeError) as exc:
|
|
60
|
+
Logger().warning(f"Shared cache disabled after put failure: {exc!r}")
|
|
61
|
+
self._backend = None
|
iaml/cache_keys.py
ADDED
|
@@ -0,0 +1,57 @@
|
|
|
1
|
+
import hashlib
|
|
2
|
+
import pickle
|
|
3
|
+
from typing import Any
|
|
4
|
+
|
|
5
|
+
import numpy as np
|
|
6
|
+
import pandas as pd
|
|
7
|
+
|
|
8
|
+
def hash_df(df: pd.DataFrame) -> str:
|
|
9
|
+
h = hashlib.sha256()
|
|
10
|
+
|
|
11
|
+
vals = pd.util.hash_pandas_object(df, index=False, categorize=True).values.tobytes()
|
|
12
|
+
h.update(vals)
|
|
13
|
+
|
|
14
|
+
h.update(pd.util.hash_pandas_object(df.index, categorize=True).values.tobytes())
|
|
15
|
+
|
|
16
|
+
h.update(pd.util.hash_pandas_object(df.columns, categorize=True).values.tobytes())
|
|
17
|
+
|
|
18
|
+
dtype_idx = pd.Index([getattr(dt, "name", str(dt)) for dt in df.dtypes])
|
|
19
|
+
h.update(pd.util.hash_pandas_object(dtype_idx).values.tobytes())
|
|
20
|
+
|
|
21
|
+
return h.hexdigest()
|
|
22
|
+
|
|
23
|
+
|
|
24
|
+
def hash_dataset(
|
|
25
|
+
X: pd.DataFrame,
|
|
26
|
+
y: Any,
|
|
27
|
+
groups: pd.DataFrame | None = None,
|
|
28
|
+
columns_types: dict | None = None,
|
|
29
|
+
target_type: str | None = None,
|
|
30
|
+
) -> str:
|
|
31
|
+
"""Hash all inputs that can affect a supervised step or evaluation.
|
|
32
|
+
|
|
33
|
+
Targets are positional, as in Dataset, so their pandas index is not used.
|
|
34
|
+
Serialization preserves their shape and dtype, including structured survival
|
|
35
|
+
targets. This only serializes local inputs; no pickle is loaded here.
|
|
36
|
+
"""
|
|
37
|
+
payload = (
|
|
38
|
+
"iaml-dataset-v2",
|
|
39
|
+
hash_df(X),
|
|
40
|
+
np.asarray(y),
|
|
41
|
+
hash_df(groups) if groups is not None else None,
|
|
42
|
+
columns_types,
|
|
43
|
+
target_type,
|
|
44
|
+
)
|
|
45
|
+
return hashlib.sha256(pickle.dumps(payload, protocol=5)).hexdigest()
|
|
46
|
+
|
|
47
|
+
|
|
48
|
+
def hash_evaluation_context(*values: Any) -> str | None:
|
|
49
|
+
"""Identify serializable evaluation settings, or decline to cache them.
|
|
50
|
+
|
|
51
|
+
Partial splitters include their arguments. Local functions and lambdas that
|
|
52
|
+
cannot be serialized are evaluated without partition or score caching.
|
|
53
|
+
"""
|
|
54
|
+
try:
|
|
55
|
+
return hashlib.sha256(pickle.dumps(values, protocol=5)).hexdigest()
|
|
56
|
+
except (pickle.PicklingError, TypeError, AttributeError, ValueError):
|
|
57
|
+
return None
|