PyIAML 1.0.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- iaml/__init__.py +56 -0
- iaml/actionable.py +11 -0
- iaml/actionables/__init__.py +21 -0
- iaml/actionables/boosting/__init__.py +4 -0
- iaml/actionables/boosting/act_adaboost.py +59 -0
- iaml/actionables/cleaning/__init__.py +26 -0
- iaml/actionables/cleaning/act_categorical_imputer.py +124 -0
- iaml/actionables/cleaning/act_count_vectorizer.py +204 -0
- iaml/actionables/cleaning/act_drop_categorical_column.py +51 -0
- iaml/actionables/cleaning/act_drop_date_column.py +48 -0
- iaml/actionables/cleaning/act_drop_high_cardinality_categorical.py +337 -0
- iaml/actionables/cleaning/act_drop_numerical_column.py +75 -0
- iaml/actionables/cleaning/act_drop_textual_column.py +51 -0
- iaml/actionables/cleaning/act_encode_target_column.py +56 -0
- iaml/actionables/cleaning/act_frequency_encoder.py +127 -0
- iaml/actionables/cleaning/act_hashing_vectorizer.py +186 -0
- iaml/actionables/cleaning/act_knn_imputer.py +152 -0
- iaml/actionables/cleaning/act_mean_column.py +79 -0
- iaml/actionables/cleaning/act_mice.py +464 -0
- iaml/actionables/cleaning/act_missing_count_feature.py +109 -0
- iaml/actionables/cleaning/act_missing_indicator.py +124 -0
- iaml/actionables/cleaning/act_onehot.py +65 -0
- iaml/actionables/cleaning/act_ordinal_encoder.py +177 -0
- iaml/actionables/cleaning/act_rare_category_grouper.py +173 -0
- iaml/actionables/cleaning/act_simple_imputer.py +109 -0
- iaml/actionables/cleaning/act_split_date.py +68 -0
- iaml/actionables/cleaning/act_target_encoder.py +274 -0
- iaml/actionables/cleaning/act_text_normalizer.py +241 -0
- iaml/actionables/cleaning/act_tf_idf.py +80 -0
- iaml/actionables/cleaning/act_word2vec.py +150 -0
- iaml/actionables/features_precleaning/__init__.py +12 -0
- iaml/actionables/features_precleaning/act_coerce_numeric_strings.py +194 -0
- iaml/actionables/features_precleaning/act_date_converter.py +99 -0
- iaml/actionables/features_precleaning/act_drop_bad_quality_rows.py +77 -0
- iaml/actionables/features_precleaning/act_drop_duplicate_rows.py +131 -0
- iaml/actionables/features_precleaning/act_drop_high_missing_columns.py +94 -0
- iaml/actionables/features_precleaning/act_drop_id_like_columns.py +294 -0
- iaml/actionables/features_precleaning/act_normalize_column_names.py +157 -0
- iaml/actionables/features_precleaning/act_sentinel_to_na_n.py +270 -0
- iaml/actionables/features_precleaning/act_trim_space.py +79 -0
- iaml/actionables/features_preprocessing/__init__.py +18 -0
- iaml/actionables/features_preprocessing/act_cyclical_date_encoding.py +212 -0
- iaml/actionables/features_preprocessing/act_fast_ica.py +161 -0
- iaml/actionables/features_preprocessing/act_feature_agglomeration.py +90 -0
- iaml/actionables/features_preprocessing/act_k_bins_discretizer.py +207 -0
- iaml/actionables/features_preprocessing/act_k_means_features.py +296 -0
- iaml/actionables/features_preprocessing/act_kernel_pca.py +143 -0
- iaml/actionables/features_preprocessing/act_log_transformer.py +122 -0
- iaml/actionables/features_preprocessing/act_nystroem.py +100 -0
- iaml/actionables/features_preprocessing/act_pca.py +77 -0
- iaml/actionables/features_preprocessing/act_polynomial_features.py +86 -0
- iaml/actionables/features_preprocessing/act_power_transformer.py +106 -0
- iaml/actionables/features_preprocessing/act_quantile_transformer.py +114 -0
- iaml/actionables/features_preprocessing/act_rbf_sampler.py +88 -0
- iaml/actionables/features_preprocessing/act_select_percentile.py +112 -0
- iaml/actionables/features_preprocessing/act_sparse_random_projection.py +157 -0
- iaml/actionables/features_preprocessing/act_truncated_svd.py +137 -0
- iaml/actionables/features_selection/__init__.py +8 -0
- iaml/actionables/features_selection/act_permutation_importance_selector.py +421 -0
- iaml/actionables/features_selection/act_remove_high_correlated_column.py +70 -0
- iaml/actionables/features_selection/act_remove_low_variance_column.py +74 -0
- iaml/actionables/features_selection/act_rfe.py +214 -0
- iaml/actionables/features_selection/act_select_from_model.py +325 -0
- iaml/actionables/features_selection/act_select_k_best.py +181 -0
- iaml/actionables/features_selection/act_vif_selector.py +130 -0
- iaml/actionables/imbalance/__init__.py +10 -0
- iaml/actionables/imbalance/act_adasyn.py +150 -0
- iaml/actionables/imbalance/act_borderline_smote.py +171 -0
- iaml/actionables/imbalance/act_near_miss.py +158 -0
- iaml/actionables/imbalance/act_random_over_sampling.py +60 -0
- iaml/actionables/imbalance/act_random_under_sampler.py +135 -0
- iaml/actionables/imbalance/act_smote.py +162 -0
- iaml/actionables/imbalance/act_smote_tomek.py +182 -0
- iaml/actionables/imbalance/act_smoteenn.py +193 -0
- iaml/actionables/imbalance/act_tomek_links.py +138 -0
- iaml/actionables/normalize/__init__.py +6 -0
- iaml/actionables/normalize/act_max_abs_scaler.py +78 -0
- iaml/actionables/normalize/act_minmax_scaler.py +56 -0
- iaml/actionables/normalize/act_normalizer.py +95 -0
- iaml/actionables/normalize/act_robust_scaler.py +111 -0
- iaml/actionables/normalize/act_standard_scaler.py +55 -0
- iaml/actionables/predictors/__init__.py +6 -0
- iaml/actionables/predictors/_xgboost.py +16 -0
- iaml/actionables/predictors/classifier/__init__.py +26 -0
- iaml/actionables/predictors/classifier/act_bagging_classifier.py +113 -0
- iaml/actionables/predictors/classifier/act_bernoulli_nb.py +89 -0
- iaml/actionables/predictors/classifier/act_catboost_classifier.py +135 -0
- iaml/actionables/predictors/classifier/act_complement_nb.py +106 -0
- iaml/actionables/predictors/classifier/act_decision_tree_classifier.py +117 -0
- iaml/actionables/predictors/classifier/act_extra_trees_classifier.py +115 -0
- iaml/actionables/predictors/classifier/act_gaussian_nb.py +53 -0
- iaml/actionables/predictors/classifier/act_hist_gradient_boosting_classifier.py +144 -0
- iaml/actionables/predictors/classifier/act_knn.py +86 -0
- iaml/actionables/predictors/classifier/act_light_gbm_classifier.py +211 -0
- iaml/actionables/predictors/classifier/act_linear_discriminant_analysis.py +63 -0
- iaml/actionables/predictors/classifier/act_linear_svc.py +134 -0
- iaml/actionables/predictors/classifier/act_logistic_regression.py +92 -0
- iaml/actionables/predictors/classifier/act_mlp_classifier.py +107 -0
- iaml/actionables/predictors/classifier/act_multinomial_nb.py +76 -0
- iaml/actionables/predictors/classifier/act_passive_aggressive_classifier.py +141 -0
- iaml/actionables/predictors/classifier/act_quadratic_discriminant_analysis.py +72 -0
- iaml/actionables/predictors/classifier/act_randomforest.py +113 -0
- iaml/actionables/predictors/classifier/act_ridge_classifier.py +116 -0
- iaml/actionables/predictors/classifier/act_sgd_classifier.py +149 -0
- iaml/actionables/predictors/classifier/act_svm_svc.py +88 -0
- iaml/actionables/predictors/classifier/act_xgboost.py +111 -0
- iaml/actionables/predictors/regressor/__init__.py +27 -0
- iaml/actionables/predictors/regressor/act_ada_boost_regressor.py +75 -0
- iaml/actionables/predictors/regressor/act_ard_regression.py +95 -0
- iaml/actionables/predictors/regressor/act_catboost_regressor.py +134 -0
- iaml/actionables/predictors/regressor/act_decision_tree_regressor.py +111 -0
- iaml/actionables/predictors/regressor/act_elastic_net_regressor.py +109 -0
- iaml/actionables/predictors/regressor/act_extra_trees_regressor.py +113 -0
- iaml/actionables/predictors/regressor/act_gaussian_process_regressor.py +55 -0
- iaml/actionables/predictors/regressor/act_gboost_regressor.py +95 -0
- iaml/actionables/predictors/regressor/act_hist_gradient_boosting_regressor.py +105 -0
- iaml/actionables/predictors/regressor/act_huber_regressor.py +101 -0
- iaml/actionables/predictors/regressor/act_knn_regressor.py +86 -0
- iaml/actionables/predictors/regressor/act_lasso_regressor.py +103 -0
- iaml/actionables/predictors/regressor/act_light_gbm_regressor.py +201 -0
- iaml/actionables/predictors/regressor/act_linear_regression.py +43 -0
- iaml/actionables/predictors/regressor/act_mlp_regressor.py +104 -0
- iaml/actionables/predictors/regressor/act_poisson_regressor.py +111 -0
- iaml/actionables/predictors/regressor/act_quantile_regressor.py +87 -0
- iaml/actionables/predictors/regressor/act_randomforest_regressor.py +116 -0
- iaml/actionables/predictors/regressor/act_ransac_regressor.py +106 -0
- iaml/actionables/predictors/regressor/act_ridge_regressor.py +107 -0
- iaml/actionables/predictors/regressor/act_sgd_regressor.py +106 -0
- iaml/actionables/predictors/regressor/act_svm_svr.py +81 -0
- iaml/actionables/predictors/regressor/act_xgboost_regressor.py +97 -0
- iaml/actionables/predictors/survival/__init__.py +12 -0
- iaml/actionables/predictors/survival/act_aalen_additive_model.py +83 -0
- iaml/actionables/predictors/survival/act_cox.py +110 -0
- iaml/actionables/predictors/survival/act_coxnet_survival_analysis.py +134 -0
- iaml/actionables/predictors/survival/act_extra_survival_trees.py +101 -0
- iaml/actionables/predictors/survival/act_fast_survival_svm.py +102 -0
- iaml/actionables/predictors/survival/act_gradient_boosting_survival_analysis.py +93 -0
- iaml/actionables/predictors/survival/act_random_survival_forest.py +91 -0
- iaml/actionables/predictors/survival/act_survival_component_wise_gboost.py +80 -0
- iaml/actionables/predictors/survival/act_survival_tree.py +120 -0
- iaml/actionables/predictors/survival/act_survival_xgboost.py +9 -0
- iaml/actionables/predictors/survival/act_weibull_aft.py +230 -0
- iaml/cache.py +61 -0
- iaml/cache_keys.py +57 -0
- iaml/candidate.py +736 -0
- iaml/core_dispatcher.py +125 -0
- iaml/data_type.py +11 -0
- iaml/dataset.py +506 -0
- iaml/decorators/__init__.py +3 -0
- iaml/decorators/all.py +4 -0
- iaml/decorators/is_step.py +45 -0
- iaml/decorators/runner.py +100 -0
- iaml/explanation.py +112 -0
- iaml/iaml.py +1072 -0
- iaml/iaml_pipeline.py +600 -0
- iaml/logger.py +138 -0
- iaml/meta_explorer_step.py +62 -0
- iaml/meta_ordered_step.py +28 -0
- iaml/meta_partial_explorer_step.py +34 -0
- iaml/meta_singleton.py +24 -0
- iaml/metastep.py +211 -0
- iaml/metric.py +111 -0
- iaml/metric_plot.py +82 -0
- iaml/metrics/__init__.py +21 -0
- iaml/metrics/_classification.py +28 -0
- iaml/metrics/_survival_times.py +22 -0
- iaml/metrics/accuracy_metric.py +59 -0
- iaml/metrics/balanced_accuracy_metric.py +67 -0
- iaml/metrics/brier_score.py +90 -0
- iaml/metrics/classification_error_metric.py +66 -0
- iaml/metrics/concordance_index_ipcw.py +84 -0
- iaml/metrics/concordance_index_metric.py +67 -0
- iaml/metrics/cumulative_dynamic_auc.py +119 -0
- iaml/metrics/f1_score_metric.py +71 -0
- iaml/metrics/integrated_brier_score.py +98 -0
- iaml/metrics/integrated_brier_score_loss.py +41 -0
- iaml/metrics/mean_absolute_error_metric.py +46 -0
- iaml/metrics/mean_squared_error_metric.py +46 -0
- iaml/metrics/mean_squared_log_error_metric.py +49 -0
- iaml/metrics/median_absolute_error_metric.py +48 -0
- iaml/metrics/precision_metric.py +63 -0
- iaml/metrics/r2_score_metric.py +45 -0
- iaml/metrics/recall_metric.py +65 -0
- iaml/metrics/roc_auc_metric.py +50 -0
- iaml/metrics/specificity_metric.py +44 -0
- iaml/metrics/specificity_multiclass_metric.py +55 -0
- iaml/metrics/specificity_multilabel_metric.py +60 -0
- iaml/optimizers/__init__.py +5 -0
- iaml/optimizers/bayesian_optimizer.py +193 -0
- iaml/optimizers/genetic_optimizer.py +284 -0
- iaml/optimizers/optimizer.py +31 -0
- iaml/optimizers/random_optimizer.py +101 -0
- iaml/plot.py +138 -0
- iaml/plots/__init__.py +32 -0
- iaml/plots/bar_plot.py +141 -0
- iaml/plots/box_plot.py +166 -0
- iaml/plots/class_prediction_error_plot.py +37 -0
- iaml/plots/classification_report_plot.py +35 -0
- iaml/plots/confusion_matrix_plot.py +34 -0
- iaml/plots/correlation_heatmap_plot.py +201 -0
- iaml/plots/cumulative_hazard_plot.py +72 -0
- iaml/plots/density_plot.py +210 -0
- iaml/plots/histogram_plot.py +179 -0
- iaml/plots/kaplan_meier_comparison_plot.py +89 -0
- iaml/plots/line_plot.py +70 -0
- iaml/plots/missingness_heatmap_plot.py +203 -0
- iaml/plots/outlier_plot.py +217 -0
- iaml/plots/pair_plot.py +228 -0
- iaml/plots/precision_recall_curve_plot.py +86 -0
- iaml/plots/prediction_error_plot.py +34 -0
- iaml/plots/qq_plot.py +220 -0
- iaml/plots/residual_plot.py +38 -0
- iaml/plots/roc_dynamique_curve_plot.py +79 -0
- iaml/plots/rocauc_plot.py +96 -0
- iaml/plots/shap_plot.py +187 -0
- iaml/plots/target_distribution_plot.py +241 -0
- iaml/plots/violin_plot.py +206 -0
- iaml/predictor.py +139 -0
- iaml/reference.py +65 -0
- iaml/shared_cache.py +90 -0
- iaml/sklearn_preprocessor.py +74 -0
- iaml/splitters/__init__.py +3 -0
- iaml/splitters/kfold_splitter.py +32 -0
- iaml/splitters/random_splitter.py +26 -0
- iaml/stack.py +39 -0
- iaml/statistic.py +66 -0
- iaml/statistics/__init__.py +77 -0
- iaml/statistics/anova_statistic.py +80 -0
- iaml/statistics/cardinality_ratio_statistic.py +63 -0
- iaml/statistics/category_cooccurrence_statistic.py +79 -0
- iaml/statistics/chi_square_statistic.py +81 -0
- iaml/statistics/coef_variation_statistic.py +72 -0
- iaml/statistics/correlation_with_target.py +105 -0
- iaml/statistics/count.py +72 -0
- iaml/statistics/data_type_summary_statistic.py +74 -0
- iaml/statistics/duplicate_row_statistic.py +56 -0
- iaml/statistics/effect_size_statistic.py +129 -0
- iaml/statistics/entropy_statistic.py +69 -0
- iaml/statistics/event_rate_statistic.py +52 -0
- iaml/statistics/grouped_mean_statistic.py +60 -0
- iaml/statistics/iqr_statistic.py +66 -0
- iaml/statistics/kurtosis.py +50 -0
- iaml/statistics/mad_statistic.py +66 -0
- iaml/statistics/mean.py +61 -0
- iaml/statistics/median_statistic.py +61 -0
- iaml/statistics/minmax.py +60 -0
- iaml/statistics/missing_rate_statistic.py +62 -0
- iaml/statistics/mode.py +47 -0
- iaml/statistics/most_frequent_ratio.py +81 -0
- iaml/statistics/outlier_count_iqr_statistic.py +76 -0
- iaml/statistics/quantile.py +59 -0
- iaml/statistics/range.py +53 -0
- iaml/statistics/rare_category_rate.py +92 -0
- iaml/statistics/skewness.py +53 -0
- iaml/statistics/stdev.py +50 -0
- iaml/statistics/summary_table_statistic.py +60 -0
- iaml/statistics/time_by_group_statistic.py +83 -0
- iaml/statistics/time_summary_statistic.py +56 -0
- iaml/statistics/top_k_value_counts.py +68 -0
- iaml/statistics/unique_count_statistic.py +57 -0
- iaml/statistics/value_counts.py +63 -0
- iaml/statistics/variance.py +51 -0
- iaml/statistics/violin.py +63 -0
- iaml/step.py +600 -0
- iaml/step_cache.py +87 -0
- iaml/step_wrapper.py +79 -0
- iaml/timed_pool_executor.py +492 -0
- iaml/type_of_target.py +68 -0
- iaml/void_step.py +101 -0
- iaml/worker_manager.py +169 -0
- iaml/wrapper/__init__.py +4 -0
- iaml/wrapper/wrap_basic_gridsearch.py +68 -0
- iaml/wrapper/wrap_genetic_gridsearch.py +293 -0
- iaml/wrapper/wrap_iterative_gridsearch.py +399 -0
- pyiaml-1.0.0.dist-info/METADATA +802 -0
- pyiaml-1.0.0.dist-info/RECORD +279 -0
- pyiaml-1.0.0.dist-info/WHEEL +5 -0
- pyiaml-1.0.0.dist-info/licenses/LICENSE +674 -0
- pyiaml-1.0.0.dist-info/top_level.txt +1 -0
iaml/metrics/__init__.py
ADDED
|
@@ -0,0 +1,21 @@
|
|
|
1
|
+
"""All metrics to evaluate models"""
|
|
2
|
+
from .accuracy_metric import AccuracyMetric
|
|
3
|
+
from .balanced_accuracy_metric import BalancedAccuracyMetric
|
|
4
|
+
from .brier_score import BrierScoreMetric
|
|
5
|
+
from .classification_error_metric import ClassificationErrorMetric
|
|
6
|
+
from .concordance_index_ipcw import ConcordanceIndexIPCWMetric
|
|
7
|
+
from .concordance_index_metric import ConcordanceIndexMetric
|
|
8
|
+
from .f1_score_metric import F1ScoreMetric
|
|
9
|
+
from .integrated_brier_score_loss import IntegratedBrierScoreLossMetric
|
|
10
|
+
from .integrated_brier_score import IntegratedBrierScoreMetric
|
|
11
|
+
from .mean_absolute_error_metric import MeanAbsoluteErrorMetric
|
|
12
|
+
from .mean_squared_error_metric import MeanSquaredErrorMetric
|
|
13
|
+
from .mean_squared_log_error_metric import MeanSquaredLogErrorMetric
|
|
14
|
+
from .median_absolute_error_metric import MedianAbsoluteErrorMetric
|
|
15
|
+
from .precision_metric import PrecisionMetric
|
|
16
|
+
from .r2_score_metric import R2ScoreMetric
|
|
17
|
+
from .recall_metric import RecallMetric
|
|
18
|
+
from .specificity_metric import SpecificityMetric
|
|
19
|
+
from .specificity_multiclass_metric import SpecificityMulticlassMetric
|
|
20
|
+
from .specificity_multilabel_metric import SpecificityMultilabelMetric
|
|
21
|
+
from .roc_auc_metric import RocAucMetric
|
|
@@ -0,0 +1,28 @@
|
|
|
1
|
+
"""Shared positive-class convention for binary classification metrics."""
|
|
2
|
+
from typing import Any
|
|
3
|
+
|
|
4
|
+
from sklearn.utils.multiclass import unique_labels
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
def resolve_pos_label(y: Any, pos_label: Any = None, y_train: Any = None) -> Any:
|
|
8
|
+
"""Resolve the positive label without depending on row order or predictions.
|
|
9
|
+
|
|
10
|
+
Training labels take precedence when available. For 0/1 and -1/1 targets,
|
|
11
|
+
the positive label is always 1, including all-negative evaluation subsets.
|
|
12
|
+
Otherwise, use the last of two sorted labels. A single nonstandard label
|
|
13
|
+
is ambiguous and requires an explicit positive label or both training classes.
|
|
14
|
+
"""
|
|
15
|
+
if pos_label is not None:
|
|
16
|
+
return pos_label
|
|
17
|
+
|
|
18
|
+
labels = unique_labels(y_train if y_train is not None else y)
|
|
19
|
+
if 0 < len(labels) <= 2:
|
|
20
|
+
if all(label in (0, 1) for label in labels) or all(label in (-1, 1) for label in labels):
|
|
21
|
+
return 1
|
|
22
|
+
if len(labels) == 2:
|
|
23
|
+
return labels[-1]
|
|
24
|
+
|
|
25
|
+
raise ValueError(
|
|
26
|
+
"Cannot infer the positive class: set pos_label explicitly or provide "
|
|
27
|
+
"y_train containing both binary classes."
|
|
28
|
+
)
|
|
@@ -0,0 +1,22 @@
|
|
|
1
|
+
"""Evaluation times shared by the survival Brier metrics."""
|
|
2
|
+
import numpy as np
|
|
3
|
+
|
|
4
|
+
|
|
5
|
+
def brier_evaluation_times(durations: np.ndarray) -> np.ndarray:
|
|
6
|
+
"""Return up to 100 evenly spaced times in [min(durations), max(durations)).
|
|
7
|
+
|
|
8
|
+
A relative grid is independent of the time unit. The endpoint is excluded
|
|
9
|
+
to keep evaluations inside the follow-up interval, including when the last
|
|
10
|
+
observation is censored.
|
|
11
|
+
"""
|
|
12
|
+
durations = np.asarray(durations, dtype=float)
|
|
13
|
+
if durations.size == 0 or not np.isfinite(durations).all():
|
|
14
|
+
raise ValueError("Brier scores require non-empty, finite follow-up times.")
|
|
15
|
+
|
|
16
|
+
start, stop = durations.min(), durations.max()
|
|
17
|
+
if start >= stop:
|
|
18
|
+
raise ValueError("Brier scores require at least two distinct follow-up times.")
|
|
19
|
+
|
|
20
|
+
times = np.linspace(start, stop, num=100, endpoint=False)
|
|
21
|
+
# Very narrow intervals can produce duplicates or round up to the endpoint.
|
|
22
|
+
return np.unique(times[times < stop])
|
|
@@ -0,0 +1,59 @@
|
|
|
1
|
+
"""[METRIC] Accuracy"""
|
|
2
|
+
from typing import Any
|
|
3
|
+
import textwrap
|
|
4
|
+
from collections import Counter
|
|
5
|
+
from sklearn.metrics import accuracy_score
|
|
6
|
+
import pandas as pd
|
|
7
|
+
from numpy import ndarray
|
|
8
|
+
from ..metric import Metric
|
|
9
|
+
|
|
10
|
+
class AccuracyMetric(Metric):
|
|
11
|
+
"""[METRIC] Accuracy"""
|
|
12
|
+
name: str = 'Accuracy'
|
|
13
|
+
_description: str = textwrap.dedent('''\
|
|
14
|
+
Accuracy measures how well a model predicts outcomes by calculating the percentage
|
|
15
|
+
of correct predictions out of the total predictions. Higher accuracy means better performance.
|
|
16
|
+
''')
|
|
17
|
+
_description_long: str = textwrap.dedent('''\
|
|
18
|
+
Accuracy is a tool to evaluate how well a predictive
|
|
19
|
+
model works, especially in healthcare.
|
|
20
|
+
It shows the percentage of correct predictions made by the model.
|
|
21
|
+
To calculate it, you add the number of correct positive and negative predictions,
|
|
22
|
+
then divide by the total number of predictions.
|
|
23
|
+
For example, if a model is correct 80 times out of 100, its accuracy is 80%.
|
|
24
|
+
''')
|
|
25
|
+
refs: list[dict[str, Any]] = [
|
|
26
|
+
{
|
|
27
|
+
'year': 2006,
|
|
28
|
+
'name': 'Understanding the meaning of accuracy, trueness and precision',
|
|
29
|
+
'authors': [
|
|
30
|
+
'Antonio Menditto',
|
|
31
|
+
'Marina Patriarca',
|
|
32
|
+
'Bertil Magnusson'
|
|
33
|
+
],
|
|
34
|
+
'doi': 'https://doi.org/10.1007/s00769-006-0191-z',
|
|
35
|
+
'publisher': ' Accreditation and Quality Assurance, Volume 12, pages 45--47'
|
|
36
|
+
}
|
|
37
|
+
]
|
|
38
|
+
|
|
39
|
+
def __str__(self) -> str:
|
|
40
|
+
return 'accuracy'
|
|
41
|
+
|
|
42
|
+
def __is_balanced(self, y: ndarray | pd.Series) -> bool:
|
|
43
|
+
"""Get one label (numpy.array or pd.series) and
|
|
44
|
+
return true if classes is balanced
|
|
45
|
+
|
|
46
|
+
:return: Whether the dataset is balanced.
|
|
47
|
+
"""
|
|
48
|
+
class_count = Counter(y)
|
|
49
|
+
total_samples = y.shape[0]
|
|
50
|
+
ideal_count = total_samples/len(class_count)
|
|
51
|
+
threshold = 0.20 * ideal_count
|
|
52
|
+
|
|
53
|
+
return not any(abs(count - ideal_count) > threshold for count in class_count.values())
|
|
54
|
+
|
|
55
|
+
def suitable(self, X: pd.DataFrame, y: pd.DataFrame, type_of_target: str) -> bool:
|
|
56
|
+
return type_of_target in ['binary', 'multiclass'] and not self.__is_balanced(y)
|
|
57
|
+
|
|
58
|
+
def compute(self, y: pd.DataFrame, y_pred: pd.DataFrame, **kwargs) -> float:
|
|
59
|
+
return accuracy_score(y, y_pred)
|
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
"""[METRIC] Balanced Accuracy"""
|
|
2
|
+
from typing import Any
|
|
3
|
+
import textwrap
|
|
4
|
+
from sklearn.metrics import balanced_accuracy_score
|
|
5
|
+
import pandas as pd
|
|
6
|
+
from ..metric import Metric
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class BalancedAccuracyMetric(Metric):
|
|
10
|
+
"""[METRIC] Balanced Accuracy"""
|
|
11
|
+
|
|
12
|
+
name: str = 'Balanced Accuracy'
|
|
13
|
+
_description: str = textwrap.dedent('''\
|
|
14
|
+
Balanced Accuracy Score is a metric that evaluates a model's
|
|
15
|
+
performance by considering both positive and negative classes equally.
|
|
16
|
+
It calculates the average accuracy for each class, making it useful for imbalanced dataset.
|
|
17
|
+
''')
|
|
18
|
+
_description_long: str = textwrap.dedent('''\
|
|
19
|
+
Balanced Accuracy Score evaluates how well a predictive
|
|
20
|
+
model performs, giving equal importance to both positive and negative classes.
|
|
21
|
+
This is important in healthcare when data is imbalanced.
|
|
22
|
+
To calculate it, you find the accuracy for each class and then average those values.
|
|
23
|
+
For example, if a model has 70% accuracy for positive cases and 90% for negative cases,
|
|
24
|
+
the balanced accuracy is (70% + 90%) / 2 = 80%. This metric ensures that the model is effective
|
|
25
|
+
for all classes, making it valuable for medical decision-making.
|
|
26
|
+
''')
|
|
27
|
+
refs: list[dict[str, Any]] = [
|
|
28
|
+
{
|
|
29
|
+
'year': 2010,
|
|
30
|
+
'name': 'The Balanced Accuracy and Its Posterior Distribution',
|
|
31
|
+
'authors': [
|
|
32
|
+
'Kay Henning Brodersen',
|
|
33
|
+
'Cheng Soon Ong',
|
|
34
|
+
'Klaas Enno Stephan',
|
|
35
|
+
'Joachim M. Buhmann'
|
|
36
|
+
],
|
|
37
|
+
'doi': 'https://doi.org/10.1109/ICPR.2010.764',
|
|
38
|
+
'publisher': textwrap.dedent("""\
|
|
39
|
+
Proceedings of the 20th International Conference on Pattern Recognition, 3121-24.
|
|
40
|
+
""")
|
|
41
|
+
},
|
|
42
|
+
{
|
|
43
|
+
'year': 2015,
|
|
44
|
+
'name': textwrap.dedent("""\
|
|
45
|
+
Fundamentals of Machine Learning for Predictive Data Analytics: Algorithms,
|
|
46
|
+
Worked Examples, and Case Studies.
|
|
47
|
+
"""),
|
|
48
|
+
'authors': [
|
|
49
|
+
'John D. Kelleher',
|
|
50
|
+
'Brian Mac Namee',
|
|
51
|
+
'Aoife D\'Arcy'
|
|
52
|
+
],
|
|
53
|
+
'doi': None,
|
|
54
|
+
'publisher': textwrap.dedent("""\
|
|
55
|
+
Fundamentals of Machine Learning for Predictive Data Analytics: Algorithms, Worked Examples, and Case Studies
|
|
56
|
+
""")
|
|
57
|
+
}
|
|
58
|
+
]
|
|
59
|
+
|
|
60
|
+
def __str__(self):
|
|
61
|
+
return 'balanced_accuracy'
|
|
62
|
+
|
|
63
|
+
def suitable(self, X: pd.DataFrame, y: pd.DataFrame, type_of_target: str) -> bool:
|
|
64
|
+
return type_of_target in ['binary', 'multiclass']
|
|
65
|
+
|
|
66
|
+
def compute(self, y: pd.DataFrame, y_pred: pd.DataFrame, **kwargs) -> float:
|
|
67
|
+
return balanced_accuracy_score(y, y_pred)
|
|
@@ -0,0 +1,90 @@
|
|
|
1
|
+
|
|
2
|
+
"""[METRIC] Brier Score for Survival Models"""
|
|
3
|
+
from typing import Any
|
|
4
|
+
import textwrap
|
|
5
|
+
import pandas as pd
|
|
6
|
+
import numpy as np
|
|
7
|
+
from sksurv.metrics import brier_score
|
|
8
|
+
from ._survival_times import brier_evaluation_times
|
|
9
|
+
from ..metric import Metric
|
|
10
|
+
from ..dataset import Dataset
|
|
11
|
+
|
|
12
|
+
class BrierScoreMetric(Metric):
|
|
13
|
+
"""[METRIC] Brier Score for Survival Models"""
|
|
14
|
+
name: str = 'Brier Score'
|
|
15
|
+
greater_is_better = False
|
|
16
|
+
_description: str = textwrap.dedent('''\
|
|
17
|
+
The Brier Score is a metric used to assess the accuracy of survival models,
|
|
18
|
+
which predict the likelihood of an event, such as death or disease, occurring within a specific timeframe.
|
|
19
|
+
It compares the model's probability predictions to actual outcomes, with lower scores indicating
|
|
20
|
+
better model performance.''')
|
|
21
|
+
_description_long: str = textwrap.dedent('''\
|
|
22
|
+
The Brier Score measures how well survival models predict the probability of an event happening, like survival over time.
|
|
23
|
+
It calculates the average squared differences between predicted probabilities and actual outcomes
|
|
24
|
+
(1 for an event occurring, 0 for it not occurring). The score ranges from 0 to 1, where 0 means perfect
|
|
25
|
+
predictions and 1 means completely inaccurate ones. This metric is valuable because it not only evaluates prediction accuracy
|
|
26
|
+
but also considers the uncertainty of those predictions. A lower Brier Score indicates a more reliable model,
|
|
27
|
+
making it a crucial tool for researchers and practitioners in fields like medicine, where accurate survival
|
|
28
|
+
predictions can significantly impact decision-making.''')
|
|
29
|
+
refs: list[dict[str, Any]] = [
|
|
30
|
+
{
|
|
31
|
+
'year': 1999,
|
|
32
|
+
'name': \
|
|
33
|
+
'Assessment and comparison of prognostic classification schemes for survival data',
|
|
34
|
+
'authors': [
|
|
35
|
+
'E. Graf',
|
|
36
|
+
'C. Schmoor',
|
|
37
|
+
'W. Sauerbrei',
|
|
38
|
+
'M. Schumacher'
|
|
39
|
+
],
|
|
40
|
+
'doi': 'https://doi.org/10.1002/(SICI)1097-0258(19990915/30)18:17/18%3C2529'\
|
|
41
|
+
'::AID-SIM274%3E3.0.CO;2-5',
|
|
42
|
+
'publisher': 'Statistics in Medicine, vol. 18, no. 17-18, pp. 2529–2545'
|
|
43
|
+
}
|
|
44
|
+
]
|
|
45
|
+
|
|
46
|
+
def __str__(self) -> str:
|
|
47
|
+
return 'brier_score'
|
|
48
|
+
|
|
49
|
+
@property
|
|
50
|
+
def needed_prediction(self) -> str:
|
|
51
|
+
return 'predict_survival_function'
|
|
52
|
+
|
|
53
|
+
def suitable(self, X: pd.DataFrame, y: pd.DataFrame, type_of_target: str) -> bool:
|
|
54
|
+
return type_of_target == 'survival'
|
|
55
|
+
|
|
56
|
+
def compute(
|
|
57
|
+
self,
|
|
58
|
+
y: pd.DataFrame,
|
|
59
|
+
y_pred: pd.DataFrame,
|
|
60
|
+
y_train: pd.DataFrame = None,
|
|
61
|
+
**kwargs) -> float:
|
|
62
|
+
"""Compute metric given y, y_pred and an optional y_train.
|
|
63
|
+
|
|
64
|
+
Evaluate at the last point of a 100-point evenly spaced grid over
|
|
65
|
+
[min(test time), max(test time)), after limiting follow-up to the
|
|
66
|
+
training horizon. The time unit does not determine the grid spacing.
|
|
67
|
+
|
|
68
|
+
:param pd.DataFrame y: Ground truth to compute the metric.
|
|
69
|
+
:param y_pred: Predicted survival probability functions, one per sample.
|
|
70
|
+
:param pd.DataFrame, optional y_train: Training ground truth. Default to None.
|
|
71
|
+
:param dict, optional \\**kwargs: Additional parameters
|
|
72
|
+
:return: Computed value
|
|
73
|
+
"""
|
|
74
|
+
y_train_samples = Dataset.normalize_survival_target(y_train)
|
|
75
|
+
y_samples = Dataset.fix_y_survival(y, y_train_samples)
|
|
76
|
+
|
|
77
|
+
y_train_struct = np.array(
|
|
78
|
+
y_train_samples,
|
|
79
|
+
dtype=[('event', 'bool'), ('time', 'float')]
|
|
80
|
+
)
|
|
81
|
+
y_struct = np.array(
|
|
82
|
+
y_samples,
|
|
83
|
+
dtype=[('event', 'bool'), ('time', 'float')]
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
time = brier_evaluation_times(y_struct['time'])[-1]
|
|
87
|
+
predictions = [fn(time) for fn in y_pred]
|
|
88
|
+
|
|
89
|
+
# Calculate the Brier score at the selected evaluation time.
|
|
90
|
+
return brier_score(y_train_struct, y_struct, predictions, time)[1][0]
|
|
@@ -0,0 +1,66 @@
|
|
|
1
|
+
"""[METRIC] Classification Error"""
|
|
2
|
+
from typing import Any
|
|
3
|
+
import textwrap
|
|
4
|
+
import pandas as pd
|
|
5
|
+
from .balanced_accuracy_metric import BalancedAccuracyMetric
|
|
6
|
+
from ..metric import Metric
|
|
7
|
+
|
|
8
|
+
|
|
9
|
+
class ClassificationErrorMetric(Metric):
|
|
10
|
+
"""[METRIC] Classification Error"""
|
|
11
|
+
greater_is_better = False
|
|
12
|
+
name: str = 'Classification Error'
|
|
13
|
+
_description: str = textwrap.dedent('''\
|
|
14
|
+
Classification Error measures a model's performance
|
|
15
|
+
by calculating the proportion of incorrect predictions. It is defined as 1
|
|
16
|
+
minus the Balanced Accuracy Score, making it useful for imbalanced datasets.
|
|
17
|
+
''')
|
|
18
|
+
_description_long: str = textwrap.dedent('''\
|
|
19
|
+
Classification Error evaluates how well a predictive model performs
|
|
20
|
+
by measuring the proportion of incorrect predictions. It is calculated as 1 minus the Balanced
|
|
21
|
+
Accuracy Score, which gives equal importance to both positive and negative classes.
|
|
22
|
+
To calculate it, you first determine the Balanced Accuracy Score, which averages the accuracy
|
|
23
|
+
of both classes. Then, you subtract that value from 1. For example, if the Balanced Accuracy Score is 80%,
|
|
24
|
+
the Classification Error would be 1 - 0.80 = 0.20, or 20%. This metric helps highlight the model's shortcomings,
|
|
25
|
+
making it a valuable tool for assessing performance in medical decision-making.''')
|
|
26
|
+
refs: list[dict[str, Any]] = [
|
|
27
|
+
{
|
|
28
|
+
'year': 2010,
|
|
29
|
+
'name': 'The Balanced Accuracy and Its Posterior Distribution',
|
|
30
|
+
'authors': [
|
|
31
|
+
'Kay Henning Brodersen',
|
|
32
|
+
'Cheng Soon Ong',
|
|
33
|
+
'Klaas Enno Stephan',
|
|
34
|
+
'Joachim M. Buhmann'
|
|
35
|
+
],
|
|
36
|
+
'doi': 'https://doi.org/10.1109/ICPR.2010.764',
|
|
37
|
+
'publisher': textwrap.dedent("""\
|
|
38
|
+
Proceedings of the 20th International Conference on Pattern Recognition, 3121-24.
|
|
39
|
+
""")
|
|
40
|
+
},
|
|
41
|
+
{
|
|
42
|
+
'year': 2015,
|
|
43
|
+
'name': textwrap.dedent("""\
|
|
44
|
+
Fundamentals of Machine Learning for Predictive Data Analytics: Algorithms, Worked Examples, and Case Studies
|
|
45
|
+
"""),
|
|
46
|
+
'authors': [
|
|
47
|
+
'John D. Kelleher',
|
|
48
|
+
'Brian Mac Namee',
|
|
49
|
+
'Aoife D\'Arcy'
|
|
50
|
+
],
|
|
51
|
+
'doi': None,
|
|
52
|
+
'publisher': textwrap.dedent("""\
|
|
53
|
+
Fundamentals of Machine Learning for Predictive Data Analytics: Algorithms, Worked Examples, and Case Studies
|
|
54
|
+
""")
|
|
55
|
+
}
|
|
56
|
+
]
|
|
57
|
+
|
|
58
|
+
def __str__(self) -> str:
|
|
59
|
+
return 'classification_error'
|
|
60
|
+
|
|
61
|
+
def suitable(self, X: pd.DataFrame, y: pd.DataFrame, type_of_target: str) -> bool:
|
|
62
|
+
return type_of_target in ['binary', 'multiclass', 'multilabel-indicator']
|
|
63
|
+
|
|
64
|
+
def compute(self, y: pd.DataFrame, y_pred: pd.DataFrame, **kwargs) -> float:
|
|
65
|
+
balanced_accuracy = BalancedAccuracyMetric().compute(y, y_pred)
|
|
66
|
+
return 1 - balanced_accuracy
|
|
@@ -0,0 +1,84 @@
|
|
|
1
|
+
"""[METRIC] Concordance Index with Inverse Probability of Censoring
|
|
2
|
+
Weights (IPCW) for Survival Models
|
|
3
|
+
"""
|
|
4
|
+
from typing import Any
|
|
5
|
+
import textwrap
|
|
6
|
+
import pandas as pd
|
|
7
|
+
import numpy as np
|
|
8
|
+
from sksurv.metrics import concordance_index_ipcw
|
|
9
|
+
from ..metric import Metric
|
|
10
|
+
from ..dataset import Dataset
|
|
11
|
+
|
|
12
|
+
class ConcordanceIndexIPCWMetric(Metric):
|
|
13
|
+
"""[METRIC] Concordance Index with Inverse Probability of Censoring
|
|
14
|
+
Weights (IPCW) for Survival Models
|
|
15
|
+
"""
|
|
16
|
+
name: str = 'C-Index IPC'
|
|
17
|
+
_description: str = textwrap.dedent('''\
|
|
18
|
+
The Concordance Index with Inverse Probability of Censoring Weights (IPCW)
|
|
19
|
+
evaluates survival models by measuring how well they predict the order of events,
|
|
20
|
+
like survival times, while accounting for censored data. A higher index indicates
|
|
21
|
+
better predictive accuracy.''')
|
|
22
|
+
_description_long: str = textwrap.dedent('''\
|
|
23
|
+
The Concordance Index with Inverse Probability of Censoring Weights (IPCW)
|
|
24
|
+
assesses survival models by focusing on the ranking of survival times. It addresses the issue of censored
|
|
25
|
+
data—when some outcomes are not fully observed—by applying weights based on the probability of censoring.
|
|
26
|
+
The index ranges from 0 to 1, with 0.5 indicating no predictive ability and 1 indicating perfect prediction.
|
|
27
|
+
By incorporating IPCW, this metric provides a more accurate evaluation of a model's performance, making it
|
|
28
|
+
essential for researchers in fields like medicine and epidemiology.''')
|
|
29
|
+
refs: list[dict[str, Any]] = [
|
|
30
|
+
{
|
|
31
|
+
'year': 2011,
|
|
32
|
+
'name': textwrap.dedent("""\
|
|
33
|
+
On the C-statistics for evaluating overall adequacy of risk prediction
|
|
34
|
+
procedures with censored survival data"""),
|
|
35
|
+
'authors': [
|
|
36
|
+
'Hajime Uno',
|
|
37
|
+
'Tianxi Cai',
|
|
38
|
+
'Michael J. Pencina',
|
|
39
|
+
'Ralph B. D\'Agostino',
|
|
40
|
+
'L. J. Wei'
|
|
41
|
+
],
|
|
42
|
+
'doi': 'https://doi.org/10.1002/sim.4154',
|
|
43
|
+
'publisher': 'Statistics in Medicine, vol. 18, no. 17-18, pp. 2529-2545'
|
|
44
|
+
}
|
|
45
|
+
]
|
|
46
|
+
|
|
47
|
+
def __str__(self) -> str:
|
|
48
|
+
return 'concordance_index_ipcw'
|
|
49
|
+
|
|
50
|
+
@property
|
|
51
|
+
def needed_prediction(self) -> str:
|
|
52
|
+
return 'predict'
|
|
53
|
+
|
|
54
|
+
def suitable(self, X: pd.DataFrame, y: pd.DataFrame, type_of_target: str) -> bool:
|
|
55
|
+
return type_of_target == 'survival'
|
|
56
|
+
|
|
57
|
+
def compute(
|
|
58
|
+
self,
|
|
59
|
+
y: pd.DataFrame,
|
|
60
|
+
y_pred: pd.DataFrame,
|
|
61
|
+
y_train: pd.DataFrame = None,
|
|
62
|
+
**kwargs) -> float:
|
|
63
|
+
"""Compute the Concordance Index (C-index) with IPCW using the predicted data.
|
|
64
|
+
|
|
65
|
+
:param pd.DataFrame y: Ground truth to compute the metric.
|
|
66
|
+
:param pd.DataFrame y_pred: Prediction to compute the metric.
|
|
67
|
+
:param pd.DataFrame y_train: Training ground truth. Default to None.
|
|
68
|
+
:param dict, optional \\**kwargs: Additional parameters
|
|
69
|
+
:return: Computed value
|
|
70
|
+
"""
|
|
71
|
+
y_train_samples = Dataset.normalize_survival_target(y_train)
|
|
72
|
+
y_samples = Dataset.fix_y_survival(y, y_train_samples)
|
|
73
|
+
|
|
74
|
+
y_train_struct = np.array(
|
|
75
|
+
y_train_samples,
|
|
76
|
+
dtype=[('event', 'bool'), ('time', 'float')]
|
|
77
|
+
)
|
|
78
|
+
y_struct = np.array(
|
|
79
|
+
y_samples,
|
|
80
|
+
dtype=[('event', 'bool'), ('time', 'float')]
|
|
81
|
+
)
|
|
82
|
+
|
|
83
|
+
# Calculate concordance index using sksurv function
|
|
84
|
+
return concordance_index_ipcw(y_train_struct, y_struct, y_pred)[0]
|
|
@@ -0,0 +1,67 @@
|
|
|
1
|
+
"""[METRIC] Concordance Index for Survival Models using sksurv"""
|
|
2
|
+
from typing import Any
|
|
3
|
+
import textwrap
|
|
4
|
+
import numpy as np
|
|
5
|
+
import pandas as pd
|
|
6
|
+
from sksurv.metrics import concordance_index_censored
|
|
7
|
+
from ..dataset import Dataset
|
|
8
|
+
from ..metric import Metric
|
|
9
|
+
|
|
10
|
+
class ConcordanceIndexMetric(Metric):
|
|
11
|
+
"""[METRIC] Concordance Index for Survival Models using sksurv"""
|
|
12
|
+
name: str = 'Concordance Index'
|
|
13
|
+
_description: str = textwrap.dedent('''\
|
|
14
|
+
The Concordance Index for Survival Models using sksurv measures how well
|
|
15
|
+
a survival model predicts the order of events, such as survival times. A higher index
|
|
16
|
+
value indicates better predictive accuracy.''')
|
|
17
|
+
_description_long: str = textwrap.dedent('''\
|
|
18
|
+
The Concordance Index for Survival Models using sksurv evaluates the performance of
|
|
19
|
+
survival models by assessing their ability to correctly rank individuals based on their
|
|
20
|
+
survival times. It focuses on the relative timing of events rather than exact predictions.
|
|
21
|
+
The index ranges from 0 to 1, where 0.5 indicates no predictive ability (similar to random guessing)
|
|
22
|
+
and 1 indicates perfect prediction of event order. This metric is particularly useful in survival
|
|
23
|
+
analysis, as it helps researchers and clinicians understand how well their models perform in
|
|
24
|
+
predicting outcomes, making it a valuable tool in fields like healthcare and clinical research.
|
|
25
|
+
''')
|
|
26
|
+
refs: list[dict[str, Any]] = [
|
|
27
|
+
{
|
|
28
|
+
'year': 1996,
|
|
29
|
+
'name': textwrap.dedent("""\
|
|
30
|
+
Multivariable prognostic models: issues in developing models, evaluating assumptions and adequacy, and measuring and reducing errors
|
|
31
|
+
"""),
|
|
32
|
+
'authors': [
|
|
33
|
+
'FRANK E.',
|
|
34
|
+
'HARRELL Jr.',
|
|
35
|
+
'KERRY L',
|
|
36
|
+
'LEE',
|
|
37
|
+
'DANIEL B. MARK'
|
|
38
|
+
],
|
|
39
|
+
'doi': textwrap.dedent("""\
|
|
40
|
+
https://doi.org/10.1002/(SICI)1097-0258(19960229)15:4%3C361::AID-SIM168%3E3.0.CO;2-4
|
|
41
|
+
"""),
|
|
42
|
+
'publisher': 'Statistics in Medicine, 15(4), 361-87'
|
|
43
|
+
}
|
|
44
|
+
]
|
|
45
|
+
|
|
46
|
+
def __str__(self) -> str:
|
|
47
|
+
return 'concordance_index'
|
|
48
|
+
|
|
49
|
+
def suitable(self, X: pd.DataFrame, y: pd.DataFrame, type_of_target: str) -> bool:
|
|
50
|
+
return type_of_target == 'survival'
|
|
51
|
+
|
|
52
|
+
def compute(
|
|
53
|
+
self,
|
|
54
|
+
y: pd.DataFrame,
|
|
55
|
+
y_pred: pd.DataFrame,
|
|
56
|
+
**kwargs) -> float:
|
|
57
|
+
samples = Dataset.normalize_survival_target(y)
|
|
58
|
+
if not samples:
|
|
59
|
+
raise ValueError("Survival targets are empty.")
|
|
60
|
+
|
|
61
|
+
events = np.asarray([event for event, _ in samples], dtype=bool)
|
|
62
|
+
times = np.asarray([time for _, time in samples], dtype=float)
|
|
63
|
+
|
|
64
|
+
# Calculate concordance index using sksurv function
|
|
65
|
+
result = concordance_index_censored(events, times, y_pred)
|
|
66
|
+
|
|
67
|
+
return result[0] # The first value in the result is the concordance index
|
|
@@ -0,0 +1,119 @@
|
|
|
1
|
+
"""Experimental cumulative AUC implementation for explicit manual investigation.
|
|
2
|
+
|
|
3
|
+
The current implementation passes survival probabilities where risk scores are
|
|
4
|
+
required, and its time grid and aggregation need validation. Importing this module
|
|
5
|
+
must not activate the metric in AutoML. See docs/component_status.rst.
|
|
6
|
+
"""
|
|
7
|
+
from typing import Any
|
|
8
|
+
import textwrap
|
|
9
|
+
import pandas as pd
|
|
10
|
+
import numpy as np
|
|
11
|
+
from sksurv.metrics import cumulative_dynamic_auc
|
|
12
|
+
from ..metric import Metric
|
|
13
|
+
from ..dataset import Dataset
|
|
14
|
+
|
|
15
|
+
class CumulativeDynamicAUCMetric(Metric):
|
|
16
|
+
"""[METRIC] Cumulative Dynamic AUC for Survival Models"""
|
|
17
|
+
name: str = 'Cumulative AUC'
|
|
18
|
+
_description: str = textwrap.dedent('''\
|
|
19
|
+
The Cumulative Dynamic AUC (Area Under the Curve) for Survival Models measures the accuracy of a survival model
|
|
20
|
+
in predicting the probability of an event over time. A higher AUC indicates better predictive performance.
|
|
21
|
+
''')
|
|
22
|
+
_description_long: str = textwrap.dedent('''\
|
|
23
|
+
The Cumulative Dynamic AUC for Survival Models evaluates how well a model predicts
|
|
24
|
+
the likelihood of an event, such as death or disease, at various time points. Unlike traditional AUC, which
|
|
25
|
+
assesses binary classification, the cumulative dynamic AUC accounts for time-dependent predictions
|
|
26
|
+
in survival analysis. This metric calculates the area under the curve of the time-dependent receiver
|
|
27
|
+
operating characteristic (ROC) curve, providing a comprehensive view of model performance over time.
|
|
28
|
+
Values range from 0 to 1, where 0.5 indicates no predictive ability and 1 indicates perfect prediction.
|
|
29
|
+
The Cumulative Dynamic AUC is particularly useful for researchers and clinicians in assessing the effectiveness
|
|
30
|
+
of survival models in real-world scenarios.
|
|
31
|
+
''')
|
|
32
|
+
refs: list[dict[str, Any]]=[
|
|
33
|
+
{
|
|
34
|
+
'year': 2007,
|
|
35
|
+
'name': textwrap.dedent("""\
|
|
36
|
+
Evaluating prediction rules for t-year survivors with censored regression models
|
|
37
|
+
"""),
|
|
38
|
+
'authors': [
|
|
39
|
+
'H. Uno',
|
|
40
|
+
'T. Cai.',
|
|
41
|
+
' L. Tian',
|
|
42
|
+
'L. J. Wei'
|
|
43
|
+
],
|
|
44
|
+
'doi': 'https://doi.org/10.1198/016214507000000149',
|
|
45
|
+
'publisher': 'Journal of the American Statistical Association, vol. 102, pp. 527–537'
|
|
46
|
+
},
|
|
47
|
+
{
|
|
48
|
+
'year': 2010,
|
|
49
|
+
'name': 'Estimation methods for time-dependent AUC models with survival data',
|
|
50
|
+
'authors': [
|
|
51
|
+
'H. Hung',
|
|
52
|
+
'C. T. Chiang'
|
|
53
|
+
],
|
|
54
|
+
'doi': '',
|
|
55
|
+
'publisher': 'Canadian Journal of Statistics, vol. 38, no. 1, pp. 8–26'
|
|
56
|
+
},
|
|
57
|
+
{
|
|
58
|
+
'year': 2014,
|
|
59
|
+
'name': textwrap.dedent("""\
|
|
60
|
+
Summary measure of discrimination in survival models based on cumulative/dynamic time-dependent ROC curves
|
|
61
|
+
"""),
|
|
62
|
+
'authors': [
|
|
63
|
+
'J. Lambert',
|
|
64
|
+
'S. Chevret'
|
|
65
|
+
],
|
|
66
|
+
'doi': '',
|
|
67
|
+
'publisher': 'Statistical Methods in Medical Research'
|
|
68
|
+
}
|
|
69
|
+
]
|
|
70
|
+
|
|
71
|
+
def __str__(self) -> str:
|
|
72
|
+
return 'cumulative_dynamic_auc'
|
|
73
|
+
|
|
74
|
+
@property
|
|
75
|
+
def needed_prediction(self) -> str:
|
|
76
|
+
return 'predict_survival_function'
|
|
77
|
+
|
|
78
|
+
def suitable(self, X: pd.DataFrame, y: pd.DataFrame, type_of_target: str) -> bool:
|
|
79
|
+
"""Keep this unvalidated implementation outside automatic metric selection."""
|
|
80
|
+
return False
|
|
81
|
+
|
|
82
|
+
def compute(
|
|
83
|
+
self,
|
|
84
|
+
y: pd.DataFrame,
|
|
85
|
+
y_pred: pd.DataFrame,
|
|
86
|
+
y_train: pd.DataFrame = None,
|
|
87
|
+
**kwargs) -> float:
|
|
88
|
+
"""Compute the cumulative dynamic AUC using the predicted data.
|
|
89
|
+
|
|
90
|
+
:param pd.DataFrame y: Ground truth data (duration and event status).
|
|
91
|
+
:param pd.DataFrame y_pred: Predicted risk scores or survival probabilities.
|
|
92
|
+
:param pd.DataFrame y_train: Training data (duration and event status).
|
|
93
|
+
:param dict, optional \\**kwargs: Additional parameters
|
|
94
|
+
:return: Computed value
|
|
95
|
+
"""
|
|
96
|
+
y_train_samples = Dataset.normalize_survival_target(y_train)
|
|
97
|
+
y_samples = Dataset.fix_y_survival(y, y_train_samples)
|
|
98
|
+
|
|
99
|
+
y_train_struct = np.array(
|
|
100
|
+
y_train_samples,
|
|
101
|
+
dtype=[('event', 'bool'), ('time', 'float')]
|
|
102
|
+
)
|
|
103
|
+
y_struct = np.array(
|
|
104
|
+
y_samples,
|
|
105
|
+
dtype=[('event', 'bool'), ('time', 'float')]
|
|
106
|
+
)
|
|
107
|
+
|
|
108
|
+
# Extract time from y test
|
|
109
|
+
_, time = zip(*y_samples)
|
|
110
|
+
times = np.arange(min(time), max(time))
|
|
111
|
+
|
|
112
|
+
# Extract risk score for each time point
|
|
113
|
+
predictions = np.asarray([[fn(t) for t in times] for fn in y_pred])
|
|
114
|
+
|
|
115
|
+
# Calculate cumulative dynamic AUC using sksurv function
|
|
116
|
+
all_points, _ = cumulative_dynamic_auc(y_train_struct, y_struct, predictions, times)
|
|
117
|
+
|
|
118
|
+
# Return mean AUC across time points
|
|
119
|
+
return np.nanmean(all_points)
|