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,83 @@
|
|
|
1
|
+
"""Experimental Aalen adapter, available only through an explicit module import.
|
|
2
|
+
|
|
3
|
+
Requires the optional, undeclared lifelines dependency. Parameter forwarding and
|
|
4
|
+
the time-by-sample hazard output do not implement IAML's predictor contract yet.
|
|
5
|
+
Excluded from automatic model selection; see docs/component_status.rst.
|
|
6
|
+
"""
|
|
7
|
+
import textwrap
|
|
8
|
+
from typing import Any
|
|
9
|
+
from lifelines import AalenAdditiveFitter
|
|
10
|
+
import pandas as pd
|
|
11
|
+
|
|
12
|
+
from ....predictor import Predictor
|
|
13
|
+
from ....candidate import Candidate
|
|
14
|
+
from ....dataset import Dataset
|
|
15
|
+
from ....decorators.all import is_step
|
|
16
|
+
|
|
17
|
+
|
|
18
|
+
@is_step('experimental')
|
|
19
|
+
class ActAalenAdditiveFitter(Predictor):
|
|
20
|
+
"""[STEP] Aalen's Additive Model for Survival Analysis"""
|
|
21
|
+
|
|
22
|
+
name: str = "AalenAdditiveFitter"
|
|
23
|
+
_usage: str = "Use when effects change over time and you want an additive alternative to ActCox. Applicable to tabular survival data with event/time and censoring. Avoid when hazards are time-constant or nonlinear interactions favor ActRandomSurvivalForest."
|
|
24
|
+
_description: str = textwrap.dedent('''\
|
|
25
|
+
Aalen's Additive Model is a semi-parametric survival analysis model
|
|
26
|
+
that estimates survival time as a function of covariates, using a linear combination
|
|
27
|
+
of time-varying covariate effects. The additive nature of the model allows it to
|
|
28
|
+
account for time-varying effects of covariates on the hazard function.''')
|
|
29
|
+
_description_long: str = textwrap.dedent('''\
|
|
30
|
+
Aalen's Additive Model is a flexible alternative to the Cox
|
|
31
|
+
Proportional Hazards model, providing time-varying covariate effects. The model
|
|
32
|
+
uses an additive approach to model the hazard function, making fewer assumptions
|
|
33
|
+
than proportional hazards models. It is particularly useful in situations where
|
|
34
|
+
covariate effects are expected to vary over time, and efficiently handles censored
|
|
35
|
+
data. The model estimates a baseline hazard function and additive contributions
|
|
36
|
+
of covariates, allowing for a more dynamic understanding of survival probabilities
|
|
37
|
+
over time.''')
|
|
38
|
+
refs: list[dict[str, Any]] = [
|
|
39
|
+
{
|
|
40
|
+
'year': 2001,
|
|
41
|
+
'name': 'Aalen’s Additive Model',
|
|
42
|
+
'authors': [
|
|
43
|
+
'O. Borgan',
|
|
44
|
+
'J. Aalen',
|
|
45
|
+
'H. Fekjær'
|
|
46
|
+
],
|
|
47
|
+
'doi': 'https://doi.org/10.1007/978-1-4757-3462-1_4',
|
|
48
|
+
'publisher': 'Survival and Event History Analysis, pages 109-142'
|
|
49
|
+
}
|
|
50
|
+
]
|
|
51
|
+
|
|
52
|
+
def __init__(self):
|
|
53
|
+
self.configuration: dict = {
|
|
54
|
+
'penalizer': {
|
|
55
|
+
'description': 'The penalizer controls the amount of L2 regularization.',
|
|
56
|
+
'default': 0.0,
|
|
57
|
+
'range': [0.0, 1.0],
|
|
58
|
+
'passthrough': False
|
|
59
|
+
}
|
|
60
|
+
}
|
|
61
|
+
self.model: AalenAdditiveFitter = None
|
|
62
|
+
|
|
63
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
64
|
+
self.model = AalenAdditiveFitter(
|
|
65
|
+
**self.passthrough_parameters()
|
|
66
|
+
)
|
|
67
|
+
|
|
68
|
+
y = pd.DataFrame(list(dataset.y), columns=['event', 'time'])
|
|
69
|
+
merged = dataset.X.reset_index(drop=True).join(y.reset_index(drop=True))
|
|
70
|
+
self.model.fit(merged, 'time', 'event')
|
|
71
|
+
return self
|
|
72
|
+
|
|
73
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
74
|
+
return dataset.type_of_target == 'survival'
|
|
75
|
+
|
|
76
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
77
|
+
return 0.5 # neutral
|
|
78
|
+
|
|
79
|
+
def predict(self, X: pd.DataFrame) -> list[float]:
|
|
80
|
+
results = self.model.predict_cumulative_hazard(X)
|
|
81
|
+
if hasattr(self, 'label_encoder'):
|
|
82
|
+
return self.label_encoder.inverse_transform(results)
|
|
83
|
+
return results
|
|
@@ -0,0 +1,110 @@
|
|
|
1
|
+
|
|
2
|
+
"""[STEP] Cox"""
|
|
3
|
+
import textwrap
|
|
4
|
+
from typing import Any
|
|
5
|
+
from sksurv.linear_model import CoxPHSurvivalAnalysis
|
|
6
|
+
|
|
7
|
+
from ....predictor import Predictor
|
|
8
|
+
from ....candidate import Candidate
|
|
9
|
+
from ....dataset import Dataset
|
|
10
|
+
from ....decorators.all import is_step
|
|
11
|
+
|
|
12
|
+
|
|
13
|
+
@is_step('predictor', 'tabular', 'survival')
|
|
14
|
+
class ActCox(Predictor):
|
|
15
|
+
"""[STEP] Cox"""
|
|
16
|
+
|
|
17
|
+
name: str = "CoxPHSurvivalAnalysis"
|
|
18
|
+
_description: str = textwrap.dedent('''\
|
|
19
|
+
CoxPHSurvivalAnalysis is a survival analysis algorithm
|
|
20
|
+
that estimates the effect of covariates on the likelihood of an event
|
|
21
|
+
occurring over time, using the Cox proportional hazards model.''')
|
|
22
|
+
_description_long: str = textwrap.dedent('''\
|
|
23
|
+
CoxPHSurvivalAnalysis is a survival analysis method
|
|
24
|
+
that models the relationship between multiple input features (covariates)
|
|
25
|
+
and the time until a particular event happens. The algorithm is based on
|
|
26
|
+
the Cox proportional hazards model, which assumes that the hazard or risk
|
|
27
|
+
of an event is a product of a baseline hazard and a factor that depends on
|
|
28
|
+
the covariates. This model is commonly used in medical research to study
|
|
29
|
+
how factors such as age, treatment, or health conditions influence survival
|
|
30
|
+
rates, or in engineering to predict equipment failure. Unlike many other
|
|
31
|
+
models, it doesn't predict the exact time of the event but estimates the
|
|
32
|
+
risk over time, handling cases where the event has not yet occurred (censored data).''')
|
|
33
|
+
_usage: str = "Use when you need a Cox model; compare ActCoxnetSurvivalAnalysis for regularization or ActRandomSurvivalForest for nonlinearity. Applicable to tabular survival data with censoring and proportional hazards. Avoid when hazards are non-proportional or interactions dominate."
|
|
34
|
+
refs: list[dict[str, Any]] = [
|
|
35
|
+
{
|
|
36
|
+
'year': 1972,
|
|
37
|
+
'name': 'Regression models and life tables',
|
|
38
|
+
'authors': [
|
|
39
|
+
'D. R. Cox'
|
|
40
|
+
],
|
|
41
|
+
'doi': 'https://doi.org/10.1111/j.2517-6161.1972.tb00899.x',
|
|
42
|
+
'publisher': 'Journal of the Royal Statistical Society. Series B, 34: page 187-220'
|
|
43
|
+
},
|
|
44
|
+
{
|
|
45
|
+
'year': 1974,
|
|
46
|
+
'name': 'Covariance Analysis of Censored Survival Data',
|
|
47
|
+
'authors': [
|
|
48
|
+
'N. E. Breslow'
|
|
49
|
+
],
|
|
50
|
+
'doi': 'https://doi.org/10.2307/2287816',
|
|
51
|
+
'publisher': 'Biometrics, 30: page 89-99'
|
|
52
|
+
},
|
|
53
|
+
{
|
|
54
|
+
'year': 1977,
|
|
55
|
+
'name': 'The Efficiency of Cox’s Likelihood Function for Censored Data',
|
|
56
|
+
'authors': [
|
|
57
|
+
'B. Efron'
|
|
58
|
+
],
|
|
59
|
+
'doi': 'https://doi.org/10.1007/978-0-387-75692-9_6',
|
|
60
|
+
'publisher': 'Journal of the American Statistical Association, 72: page 557-565'
|
|
61
|
+
}
|
|
62
|
+
]
|
|
63
|
+
|
|
64
|
+
def __init__(self):
|
|
65
|
+
self.configuration: dict = {
|
|
66
|
+
'alpha': {
|
|
67
|
+
'description': textwrap.dedent('''\
|
|
68
|
+
Regularization strength. Higher values specify stronger
|
|
69
|
+
regularization. alpha=0 means no regularization.'''),
|
|
70
|
+
'default': 1,
|
|
71
|
+
'range': [0, 100],
|
|
72
|
+
'passthrough': True
|
|
73
|
+
},
|
|
74
|
+
'ties': {
|
|
75
|
+
'description': textwrap.dedent('''\
|
|
76
|
+
Method for handling tied event times in the data.
|
|
77
|
+
"breslow" is the most common method.'''),
|
|
78
|
+
'default': 'breslow',
|
|
79
|
+
'categorical': ['breslow', 'efron']
|
|
80
|
+
},
|
|
81
|
+
'n_iter': {
|
|
82
|
+
'description': 'Maximum number of iterations for fitting the model.',
|
|
83
|
+
'default': 100,
|
|
84
|
+
'range': [1, 10000],
|
|
85
|
+
'passthrough': True
|
|
86
|
+
},
|
|
87
|
+
'tol': {
|
|
88
|
+
'description': textwrap.dedent('''\
|
|
89
|
+
Tolerance for stopping criteria. Determines the precision
|
|
90
|
+
of the solution.'''),
|
|
91
|
+
'default': 1e-09,
|
|
92
|
+
'range': [1e-12, 1e-03],
|
|
93
|
+
'passthrough': True
|
|
94
|
+
}
|
|
95
|
+
}
|
|
96
|
+
self.model: CoxPHSurvivalAnalysis = None
|
|
97
|
+
|
|
98
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
99
|
+
self.model = CoxPHSurvivalAnalysis(
|
|
100
|
+
**self.passthrough_parameters()
|
|
101
|
+
)
|
|
102
|
+
X, y = dataset.to_survival()
|
|
103
|
+
self.model.fit(X, y)
|
|
104
|
+
return self
|
|
105
|
+
|
|
106
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
107
|
+
return dataset.type_of_target == 'survival'
|
|
108
|
+
|
|
109
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
110
|
+
return 0.5 # neutral
|
|
@@ -0,0 +1,134 @@
|
|
|
1
|
+
"""[STEP] Coxnet Survival Analysis"""
|
|
2
|
+
import inspect
|
|
3
|
+
import textwrap
|
|
4
|
+
from typing import Any
|
|
5
|
+
from sksurv.linear_model import CoxnetSurvivalAnalysis
|
|
6
|
+
|
|
7
|
+
from ....predictor import Predictor
|
|
8
|
+
from ....candidate import Candidate
|
|
9
|
+
from ....dataset import Dataset
|
|
10
|
+
from ....data_type import DataType
|
|
11
|
+
from ....decorators.all import is_step
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@is_step('predictor', 'tabular', 'survival')
|
|
15
|
+
class ActCoxnetSurvivalAnalysis(Predictor):
|
|
16
|
+
"""[STEP] Coxnet Survival Analysis"""
|
|
17
|
+
|
|
18
|
+
name: str = "CoxnetSurvivalAnalysis"
|
|
19
|
+
_description: str = textwrap.dedent('''\
|
|
20
|
+
CoxnetSurvivalAnalysis fits a Cox proportional hazards model with
|
|
21
|
+
elastic-net regularization, combining L1 and L2 penalties to handle
|
|
22
|
+
high-dimensional survival data.''')
|
|
23
|
+
_description_long: str = textwrap.dedent('''\
|
|
24
|
+
CoxnetSurvivalAnalysis estimates a Cox proportional hazards model while
|
|
25
|
+
applying elastic-net regularization along a path of penalty strengths.
|
|
26
|
+
The L1 component encourages sparse feature selection, while the L2
|
|
27
|
+
component stabilizes coefficients when predictors are correlated. This
|
|
28
|
+
makes the model well suited for survival datasets with many variables
|
|
29
|
+
and right-censored observations.''')
|
|
30
|
+
_usage: str = "Use when you need regularized Cox with feature selection; compare ActCox or ActRandomSurvivalForest. Applicable to tabular survival data with many numeric, correlated predictors. Avoid when effects are highly nonlinear or inputs are mostly categorical."
|
|
31
|
+
refs: list[dict[str, Any]] = [
|
|
32
|
+
{
|
|
33
|
+
'year': 2011,
|
|
34
|
+
'name': "Regularization Paths for Cox's Proportional Hazards Model \
|
|
35
|
+
via Coordinate Descent",
|
|
36
|
+
'authors': [
|
|
37
|
+
'Noah Simon',
|
|
38
|
+
'Jerome Friedman',
|
|
39
|
+
'Trevor Hastie',
|
|
40
|
+
'Robert Tibshirani'
|
|
41
|
+
],
|
|
42
|
+
'doi': 'https://doi.org/10.18637/jss.v039.i05',
|
|
43
|
+
'publisher': 'Journal of Statistical Software, 39(5)'
|
|
44
|
+
},
|
|
45
|
+
{
|
|
46
|
+
'year': 2005,
|
|
47
|
+
'name': 'Regularization and Variable Selection via the Elastic Net',
|
|
48
|
+
'authors': [
|
|
49
|
+
'Hui Zou',
|
|
50
|
+
'Trevor Hastie'
|
|
51
|
+
],
|
|
52
|
+
'doi': 'https://doi.org/10.1111/j.1467-9868.2005.00503.x',
|
|
53
|
+
'publisher': 'Journal of the Royal Statistical Society Series B'
|
|
54
|
+
}
|
|
55
|
+
]
|
|
56
|
+
|
|
57
|
+
def __init__(self):
|
|
58
|
+
self.configuration: dict = {
|
|
59
|
+
'l1_ratio': {
|
|
60
|
+
'description': 'Mixing parameter between L1 and L2 penalty.',
|
|
61
|
+
'default': 0.5,
|
|
62
|
+
'range': [0.0, 1.0]
|
|
63
|
+
},
|
|
64
|
+
'n_alphas': {
|
|
65
|
+
'description': 'Number of alpha values along the regularization path.',
|
|
66
|
+
'default': 100,
|
|
67
|
+
'range': [10, 200]
|
|
68
|
+
},
|
|
69
|
+
'alpha_min_ratio': {
|
|
70
|
+
'description': textwrap.dedent('''\
|
|
71
|
+
Smallest alpha as a fraction of alpha_max for the regularization
|
|
72
|
+
path.'''),
|
|
73
|
+
'default': 0.01,
|
|
74
|
+
'range': [1e-04, 1.0]
|
|
75
|
+
},
|
|
76
|
+
'max_iter': {
|
|
77
|
+
'description': 'Maximum number of coordinate descent iterations.',
|
|
78
|
+
'default': 1000,
|
|
79
|
+
'range': [100, 100000]
|
|
80
|
+
},
|
|
81
|
+
'tol': {
|
|
82
|
+
'description': 'Stopping criterion.',
|
|
83
|
+
'default': 1e-07,
|
|
84
|
+
'range': [1e-09, 1e-03]
|
|
85
|
+
},
|
|
86
|
+
'fit_baseline_model': {
|
|
87
|
+
'description': textwrap.dedent('''\
|
|
88
|
+
Fit baseline hazard models to enable survival function
|
|
89
|
+
predictions.'''),
|
|
90
|
+
'default': False,
|
|
91
|
+
'categorical': [True, False]
|
|
92
|
+
}
|
|
93
|
+
}
|
|
94
|
+
self.model: CoxnetSurvivalAnalysis = None
|
|
95
|
+
self.columns: list[str] = []
|
|
96
|
+
|
|
97
|
+
def _select_features(self, X):
|
|
98
|
+
if self.columns and hasattr(X, 'columns'):
|
|
99
|
+
return X[self.columns]
|
|
100
|
+
return X
|
|
101
|
+
|
|
102
|
+
def _model_parameters(self) -> dict[str, Any]:
|
|
103
|
+
params = self.passthrough_parameters()
|
|
104
|
+
sig_params = inspect.signature(CoxnetSurvivalAnalysis).parameters
|
|
105
|
+
return {key: value for key, value in params.items() if key in sig_params}
|
|
106
|
+
|
|
107
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
108
|
+
self.columns = dataset.get_columns_names_by_type(DataType.NUMERIC)
|
|
109
|
+
if not self.columns:
|
|
110
|
+
self.columns = dataset.features
|
|
111
|
+
|
|
112
|
+
self.model = CoxnetSurvivalAnalysis(**self._model_parameters())
|
|
113
|
+
X, y = dataset.to_survival()
|
|
114
|
+
self.model.fit(self._select_features(X), y)
|
|
115
|
+
return self
|
|
116
|
+
|
|
117
|
+
def predict(self, X):
|
|
118
|
+
return super().predict(self._select_features(X))
|
|
119
|
+
|
|
120
|
+
def predict_survival_function(self, X):
|
|
121
|
+
return super().predict_survival_function(self._select_features(X))
|
|
122
|
+
|
|
123
|
+
def predict_cumulative_hazard_function(self, X):
|
|
124
|
+
return super().predict_cumulative_hazard_function(self._select_features(X))
|
|
125
|
+
|
|
126
|
+
def score(self, X, y=None, *args, **kwargs):
|
|
127
|
+
return self.model.score(self._select_features(X), y, *args, **kwargs)
|
|
128
|
+
|
|
129
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
130
|
+
return dataset.type_of_target == 'survival' \
|
|
131
|
+
and bool(dataset.get_columns_names_by_type(DataType.NUMERIC))
|
|
132
|
+
|
|
133
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
134
|
+
return 0.5 # neutral
|
|
@@ -0,0 +1,101 @@
|
|
|
1
|
+
"""[STEP] Extra Survival Trees"""
|
|
2
|
+
import textwrap
|
|
3
|
+
from typing import Any
|
|
4
|
+
from sksurv.ensemble import ExtraSurvivalTrees
|
|
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 ActExtraSurvivalTrees(Predictor):
|
|
14
|
+
"""[STEP] Extra Survival Trees"""
|
|
15
|
+
name: str = "ExtraSurvivalTrees"
|
|
16
|
+
_usage: str = "Use when you want a randomized tree ensemble for survival, as an alternative to ActRandomSurvivalForest. Applicable to tabular censored survival data with nonlinear feature effects. Avoid when you need proportional-hazards interpretability like ActCox or data is very small."
|
|
17
|
+
_description: str = textwrap.dedent('''\
|
|
18
|
+
ExtraSurvivalTrees is an ensemble learning method for survival
|
|
19
|
+
analysis based on extremely randomized trees. It fits multiple decision trees
|
|
20
|
+
to the data, where each tree is built from a random subset of features and
|
|
21
|
+
splits are selected randomly. This method provides more variance reduction
|
|
22
|
+
and robustness, especially useful when dealing with high-dimensional or
|
|
23
|
+
sparse data.''')
|
|
24
|
+
_description_long: str = textwrap.dedent('''\
|
|
25
|
+
ExtraSurvivalTrees is a variant of ensemble learning for survival
|
|
26
|
+
analysis that uses extremely randomized trees. In this approach, multiple trees
|
|
27
|
+
are grown by selecting random subsets of features and splitting points.
|
|
28
|
+
Compared to other tree-based methods, this randomness helps reduce overfitting
|
|
29
|
+
and increases model robustness. The method is particularly useful for survival
|
|
30
|
+
datasets that contain complex, non-linear relationships between features.
|
|
31
|
+
ExtraSurvivalTrees handles censored data and can provide interpretable models
|
|
32
|
+
for survival time predictions.''')
|
|
33
|
+
refs: list[dict[str, Any]] = [
|
|
34
|
+
{
|
|
35
|
+
'year': 2006,
|
|
36
|
+
'name': 'Extremely Randomized Trees',
|
|
37
|
+
'authors': [
|
|
38
|
+
'P. Geurts',
|
|
39
|
+
'D. Ernst',
|
|
40
|
+
'L. Wehenkel'
|
|
41
|
+
],
|
|
42
|
+
'doi': 'https://doi.org/10.1007/s10994-006-6226-1',
|
|
43
|
+
'publisher': 'Machine Learning, 63(1), 3-42'
|
|
44
|
+
}
|
|
45
|
+
]
|
|
46
|
+
|
|
47
|
+
def __init__(self):
|
|
48
|
+
self.configuration: dict = {
|
|
49
|
+
'n_estimators': {
|
|
50
|
+
'description': 'The number of trees in the forest.',
|
|
51
|
+
'default': 100,
|
|
52
|
+
'range': [1, 1000],
|
|
53
|
+
'passthrough': True
|
|
54
|
+
},
|
|
55
|
+
'max_depth': {
|
|
56
|
+
'description': 'The maximum depth of the trees.',
|
|
57
|
+
'default': None,
|
|
58
|
+
'range': [1, None],
|
|
59
|
+
'passthrough': True
|
|
60
|
+
},
|
|
61
|
+
'min_samples_split': {
|
|
62
|
+
'description': 'The minimum number of samples required to split an internal node.',
|
|
63
|
+
'default': 2,
|
|
64
|
+
'range': [2, 20],
|
|
65
|
+
'passthrough': True
|
|
66
|
+
},
|
|
67
|
+
'min_samples_leaf': {
|
|
68
|
+
'description': 'The minimum number of samples required to be at a leaf node.',
|
|
69
|
+
'default': 1,
|
|
70
|
+
'range': [1, 20],
|
|
71
|
+
'passthrough': True
|
|
72
|
+
},
|
|
73
|
+
'max_features': {
|
|
74
|
+
'description': textwrap.dedent('''\
|
|
75
|
+
The number of features to consider when looking for the
|
|
76
|
+
best split.'''),
|
|
77
|
+
'default': "sqrt",
|
|
78
|
+
'categorical': ["sqrt", "log2", None],
|
|
79
|
+
'passthrough': True
|
|
80
|
+
},
|
|
81
|
+
'random_state': {
|
|
82
|
+
'description': 'Random seed (integer or None) for the estimator.',
|
|
83
|
+
'default': None,
|
|
84
|
+
'passthrough': True
|
|
85
|
+
}
|
|
86
|
+
}
|
|
87
|
+
self.model: ExtraSurvivalTrees = None
|
|
88
|
+
|
|
89
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
90
|
+
self.model = ExtraSurvivalTrees(
|
|
91
|
+
**self.passthrough_parameters()
|
|
92
|
+
)
|
|
93
|
+
X, y = dataset.to_survival()
|
|
94
|
+
self.model.fit(X, y)
|
|
95
|
+
return self
|
|
96
|
+
|
|
97
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
98
|
+
return dataset.type_of_target == 'survival'
|
|
99
|
+
|
|
100
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
101
|
+
return 0.5 # neutral
|
|
@@ -0,0 +1,102 @@
|
|
|
1
|
+
"""[STEP] Fast Survival SVM"""
|
|
2
|
+
import inspect
|
|
3
|
+
import textwrap
|
|
4
|
+
from typing import Any
|
|
5
|
+
from sksurv.svm import FastSurvivalSVM
|
|
6
|
+
|
|
7
|
+
from ....predictor import Predictor
|
|
8
|
+
from ....candidate import Candidate
|
|
9
|
+
from ....dataset import Dataset
|
|
10
|
+
from ....data_type import DataType
|
|
11
|
+
from ....decorators.all import is_step
|
|
12
|
+
|
|
13
|
+
|
|
14
|
+
@is_step('predictor', 'tabular', 'survival')
|
|
15
|
+
class ActFastSurvivalSVM(Predictor):
|
|
16
|
+
"""[STEP] Fast Survival SVM"""
|
|
17
|
+
|
|
18
|
+
name: str = "FastSurvivalSVM"
|
|
19
|
+
_usage: str = "Use when you want a fast linear ranking model for survival risk, as a simpler alternative to ActCox. Applicable to right-censored tabular survival data with mostly numeric features. Avoid when nonlinear effects or interactions dominate; prefer ActRandomSurvivalForest."
|
|
20
|
+
_description: str = textwrap.dedent('''\
|
|
21
|
+
FastSurvivalSVM is a linear support vector machine for survival analysis
|
|
22
|
+
that learns a risk score using hinge-style ranking losses adapted to
|
|
23
|
+
censored data.''')
|
|
24
|
+
_description_long: str = textwrap.dedent('''\
|
|
25
|
+
FastSurvivalSVM optimizes a pairwise ranking objective so that samples
|
|
26
|
+
with earlier events receive higher risk scores. The loss is based on
|
|
27
|
+
hinge-style constraints adapted to right-censored observations, and a
|
|
28
|
+
rank_ratio parameter can blend ranking and regression terms. The model
|
|
29
|
+
is linear and efficient, making it suitable for larger tabular datasets
|
|
30
|
+
where fast, deterministic training is desired.''')
|
|
31
|
+
|
|
32
|
+
def __init__(self):
|
|
33
|
+
self.configuration: dict = {
|
|
34
|
+
'alpha': {
|
|
35
|
+
'description': textwrap.dedent('''\
|
|
36
|
+
Regularization strength. Higher values enforce stronger
|
|
37
|
+
regularization.'''),
|
|
38
|
+
'default': 1.0,
|
|
39
|
+
'range': [1e-04, 100.0]
|
|
40
|
+
},
|
|
41
|
+
'rank_ratio': {
|
|
42
|
+
'description': textwrap.dedent('''\
|
|
43
|
+
Weighting between ranking and regression losses. 1.0 uses
|
|
44
|
+
pure ranking; 0.0 uses pure regression.'''),
|
|
45
|
+
'default': 1.0,
|
|
46
|
+
'range': [0.0, 1.0]
|
|
47
|
+
},
|
|
48
|
+
'fit_intercept': {
|
|
49
|
+
'description': 'Whether to fit the intercept term.',
|
|
50
|
+
'default': True,
|
|
51
|
+
'categorical': [True, False]
|
|
52
|
+
},
|
|
53
|
+
'max_iter': {
|
|
54
|
+
'description': 'Maximum number of iterations for the optimizer.',
|
|
55
|
+
'default': 200,
|
|
56
|
+
'range': [10, 5000]
|
|
57
|
+
},
|
|
58
|
+
'tol': {
|
|
59
|
+
'description': 'Stopping tolerance.',
|
|
60
|
+
'default': 1e-05,
|
|
61
|
+
'range': [1e-08, 1e-02]
|
|
62
|
+
},
|
|
63
|
+
'random_state': {
|
|
64
|
+
'description': 'Random state for reproducibility.',
|
|
65
|
+
'default': 42
|
|
66
|
+
}
|
|
67
|
+
}
|
|
68
|
+
self.model: FastSurvivalSVM = None
|
|
69
|
+
self.columns: list[str] = []
|
|
70
|
+
|
|
71
|
+
def _select_features(self, X):
|
|
72
|
+
if self.columns and hasattr(X, 'columns'):
|
|
73
|
+
return X[self.columns]
|
|
74
|
+
return X
|
|
75
|
+
|
|
76
|
+
def _model_parameters(self) -> dict[str, Any]:
|
|
77
|
+
params = self.passthrough_parameters()
|
|
78
|
+
sig_params = inspect.signature(FastSurvivalSVM).parameters
|
|
79
|
+
return {key: value for key, value in params.items() if key in sig_params}
|
|
80
|
+
|
|
81
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
82
|
+
self.columns = dataset.get_columns_names_by_type(DataType.NUMERIC)
|
|
83
|
+
if not self.columns:
|
|
84
|
+
self.columns = dataset.features
|
|
85
|
+
|
|
86
|
+
self.model = FastSurvivalSVM(**self._model_parameters())
|
|
87
|
+
X, y = dataset.to_survival()
|
|
88
|
+
self.model.fit(self._select_features(X), y)
|
|
89
|
+
return self
|
|
90
|
+
|
|
91
|
+
def predict(self, X):
|
|
92
|
+
return super().predict(self._select_features(X))
|
|
93
|
+
|
|
94
|
+
def score(self, X, y=None, *args, **kwargs):
|
|
95
|
+
return self.model.score(self._select_features(X), y, *args, **kwargs)
|
|
96
|
+
|
|
97
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
98
|
+
return dataset.type_of_target == 'survival' \
|
|
99
|
+
and bool(dataset.get_columns_names_by_type(DataType.NUMERIC))
|
|
100
|
+
|
|
101
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
102
|
+
return 0.5 # neutral
|
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
"""[STEP] Gradient Boosting Survival Analysis"""
|
|
2
|
+
import textwrap
|
|
3
|
+
from typing import Any
|
|
4
|
+
from sksurv.ensemble import GradientBoostingSurvivalAnalysis
|
|
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 ActGradientBoostingSurvivalAnalysis(Predictor):
|
|
14
|
+
"""[STEP] Gradient Boosting Survival Analysis"""
|
|
15
|
+
|
|
16
|
+
name: str = "GradientBoostingSurvivalAnalysis"
|
|
17
|
+
_usage: str = "Use when you need non-linear survival modeling and ActCox underfits. Applicable to tabular time-to-event data with right-censoring. Avoid when you need simpler baselines or strong ensembles like ActRandomSurvivalForest."
|
|
18
|
+
_description: str = textwrap.dedent('''\
|
|
19
|
+
GradientBoostingSurvivalAnalysis is a survival analysis algorithm
|
|
20
|
+
that uses gradient boosting to model the risk of an event over time.
|
|
21
|
+
It fits an ensemble of regression trees to capture non-linear effects
|
|
22
|
+
and interactions in censored survival data.''')
|
|
23
|
+
_description_long: str = textwrap.dedent('''\
|
|
24
|
+
GradientBoostingSurvivalAnalysis extends gradient boosting to
|
|
25
|
+
time-to-event data by optimizing a survival-specific loss function.
|
|
26
|
+
The model builds an ensemble of shallow regression trees, each correcting
|
|
27
|
+
the errors of the previous ones, resulting in a flexible estimator for
|
|
28
|
+
complex covariate effects. It can handle right-censored observations and
|
|
29
|
+
is useful when proportional hazards assumptions are too restrictive.''')
|
|
30
|
+
refs: list[dict[str, Any]] = [
|
|
31
|
+
{
|
|
32
|
+
'year': 2010,
|
|
33
|
+
'name': 'Gradient boosting for survival analysis',
|
|
34
|
+
'authors': [
|
|
35
|
+
'Chen, Yifei',
|
|
36
|
+
'Jia, Zhenyu',
|
|
37
|
+
'Mercola, Dan',
|
|
38
|
+
'Xie, Xiaohui'
|
|
39
|
+
],
|
|
40
|
+
'doi': 'https://doi.org/10.1155/2013/873595',
|
|
41
|
+
'publisher': 'Advances in Data Analysis, Data Handling and Business Intelligence, \
|
|
42
|
+
pages 239-248'
|
|
43
|
+
}
|
|
44
|
+
]
|
|
45
|
+
|
|
46
|
+
def __init__(self):
|
|
47
|
+
self.configuration: dict = {
|
|
48
|
+
'n_estimators': {
|
|
49
|
+
'description': 'Number of boosting stages to be run.',
|
|
50
|
+
'default': 100,
|
|
51
|
+
'range': [1, 1000],
|
|
52
|
+
'passthrough': True
|
|
53
|
+
},
|
|
54
|
+
'learning_rate': {
|
|
55
|
+
'description': 'Learning rate shrinks the contribution of each tree by this value.',
|
|
56
|
+
'default': 0.1,
|
|
57
|
+
'range': [0.01, 1.0],
|
|
58
|
+
'passthrough': True
|
|
59
|
+
},
|
|
60
|
+
'max_depth': {
|
|
61
|
+
'description': 'The maximum depth of the individual trees.',
|
|
62
|
+
'default': 3,
|
|
63
|
+
'range': [1, 20],
|
|
64
|
+
'passthrough': True
|
|
65
|
+
},
|
|
66
|
+
'min_samples_split': {
|
|
67
|
+
'description': 'The minimum number of samples required to split an internal node.',
|
|
68
|
+
'default': 2,
|
|
69
|
+
'range': [2, 20],
|
|
70
|
+
'passthrough': True
|
|
71
|
+
},
|
|
72
|
+
'min_samples_leaf': {
|
|
73
|
+
'description': 'The minimum number of samples required to be at a leaf node.',
|
|
74
|
+
'default': 1,
|
|
75
|
+
'range': [1, 20],
|
|
76
|
+
'passthrough': True
|
|
77
|
+
}
|
|
78
|
+
}
|
|
79
|
+
self.model: GradientBoostingSurvivalAnalysis = None
|
|
80
|
+
|
|
81
|
+
def fit(self, dataset: Dataset): # pylint: disable=unused-argument
|
|
82
|
+
self.model = GradientBoostingSurvivalAnalysis(
|
|
83
|
+
**self.passthrough_parameters()
|
|
84
|
+
)
|
|
85
|
+
X, y = dataset.to_survival()
|
|
86
|
+
self.model.fit(X, y)
|
|
87
|
+
return self
|
|
88
|
+
|
|
89
|
+
def suitable(self, dataset: Dataset) -> bool:
|
|
90
|
+
return dataset.type_of_target == 'survival'
|
|
91
|
+
|
|
92
|
+
def priorize(self, candidate: Candidate = None) -> float:
|
|
93
|
+
return 0.5 # neutral
|