
Train and evaluate machine learning models for classification or survival analysis
compute_features.ML.RdThis function trains and evaluates machine learning models using cross-validation on training data
and then evaluates performance on independent test data. It supports both classification and
survival analysis tasks, including hyperparameter tuning and cohort-based
(Leave-One-Dataset-Out, LODO) validation. For survival models, it computes the C-index
and, with return = TRUE, generates Kaplan-Meier plots stratified by predicted risk.
Usage
compute_features.ML(
features_train,
features_test,
coldata,
task_type = c("classification", "survival"),
trait = NULL,
trait.positive = NULL,
time_var = NULL,
event_var = NULL,
metric = "AUROC",
k_folds = 10,
n_rep = 5,
LODO = FALSE,
batch_id = NULL,
file_name = NULL,
ncores = NULL,
return = FALSE,
fold_construction_fun = NULL,
fold_construction_args_fixed = NULL,
fold_construction_args_tunable = NULL,
seed = 123,
preprocess = TRUE
)Arguments
- features_train
A data frame or matrix of predictor variables used for training (rows = samples, columns = features).
- features_test
A data frame or matrix of predictor variables used for testing.
- coldata
A data frame containing outcome information. Its row names must include those of
features_trainandfeatures_test: the outcome of each sample is taken by row name.- task_type
Character. Type of task:
"classification"or"survival".- trait
Character. Column name in
coldataused as the target variable (required for classification tasks).- trait.positive
Value in
traitthat represents the positive class (classification only). Ensures all performance metrics and interpretability analyses consistently treat the correct class as positive.- time_var
Character. Column name in
coldatacontaining survival/follow-up time (required for survival tasks).- event_var
Character. Column name in
coldataindicating event occurrence (1 = event occurred, 0 = censored; required for survival tasks).- metric
Character. Performance metric used for model tuning and selection:
Classification:
"AUROC"(default) or"AUPRC".Survival: evaluated using concordance index (C-index).
- k_folds
Integer. Number of folds for cross-validation. Default: 10.
- n_rep
Integer. Number of repetitions for cross-validation. Default: 5.
- LODO
Logical. If
TRUE, the cross-validation folds are stratified by cohort and outcome (seecompute_features.training.ML()).- batch_id
Character. Column name in
coldatawith the cohort/batch of each sample (required ifLODO = TRUE). The cross-validation folds are then stratified by cohort and outcome.- file_name
Character. Base name used to save plots/results under
Results/. For survival tasks, Kaplan-Meier plots are saved as"Results/Survival_KM_<file_name>.pdf".- ncores
Integer. Number of CPU cores for parallelization (cross-validation folds are processed in parallel). Default:
NULL(sequential). For classification withfold_construction_fun, the models are trained sequentially andncoresis not used.- return
Logical. Whether to save the plots in
Results/. Default:FALSE.- fold_construction_fun
Function. Optional custom function to construct cross-validation folds. Must accept a
bestuneargument internally to inject optimized hyperparameters. Used for both classification and survival.features_testis used as given: it must contain the features built by this function (e.g. the test samples projected onto the structure learned on the training samples).- fold_construction_args_fixed
List. Fixed arguments passed to
fold_construction_funfor both CV and final training.- fold_construction_args_tunable
List. Arguments passed to
fold_construction_fundefining hyperparameters to explore during CV.- seed
Integer. Random seed for reproducible cross-validation (fold assignment and model fitting, including parallel runs). Default:
123. UseNULLto leave the random number generator untouched. Seecompute_features.training.ML()for details.- preprocess
Logical. If
TRUE(default), near-constant and highly correlated (|r| > 0.9) features are removed from the training features before training. UseFALSEto train on the features as given. Seecompute_features.training.ML()for details.
Value
A named list:
- Model
The output of
compute_features.training.ML()on the training set (the selected model is$Model$Model).- AUC
Classification: AUROC and AUPRC on the test set, with bootstrap confidence intervals (see
compute_prediction()).- Metrics
Classification: threshold-based performance metrics on the test set.
- Prediction
Predicted class probabilities (classification) or risk scores (survival) of the test samples.
- Curve_bands
Classification: pointwise 95% bootstrap bands around the ROC and precision-recall curves.
- C_index
Survival: C-index on the test set.
Details
For classification tasks, the function performs repeated k-fold cross-validation with hyperparameter tuning, followed by evaluation on the test set. ROC and PR curves are generated.
For survival tasks, it performs model selection using the C-index, refits the best model
on the full training data and evaluates the C-index on the test set. With return = TRUE, Kaplan-Meier
curves of the test samples split at the median predicted risk are saved, with the C-index and log-rank test
p-value.
Examples
if (FALSE) { # \dontrun{
# --- Classification ---
data(data_example_classification)
X <- data_example_classification[, setdiff(colnames(data_example_classification), "target")]
set.seed(123)
train_idx <- caret::createDataPartition(data_example_classification$target, p = 0.7, list = FALSE)
res <- compute_features.ML(features_train = X[train_idx, ],
features_test = X[-train_idx, ],
coldata = data_example_classification,
task_type = "classification",
trait = "target",
trait.positive = "1",
k_folds = 5,
n_rep = 2,
ncores = 2)
res$AUC
# --- Survival ---
data(data_example_survival)
X <- data_example_survival[, setdiff(colnames(data_example_survival), c("time", "status"))]
set.seed(123)
train_idx <- caret::createDataPartition(data_example_survival$status, p = 0.7, list = FALSE)
res_survival <- compute_features.ML(features_train = X[train_idx, ],
features_test = X[-train_idx, ],
coldata = data_example_survival,
task_type = "survival",
time_var = "time",
event_var = "status",
k_folds = 5,
n_rep = 2,
ncores = 2)
res_survival$C_index
} # }