
Compute Prediction Metrics for a Trained Machine Learning Model
compute_prediction.RdApplies a trained model to a test set and evaluates it. For classification, it computes AUROC and AUPRC with bootstrap confidence intervals, and Accuracy, Sensitivity, Specificity, Precision, Recall, F1 score and MCC at each probability threshold. For survival, it predicts risk scores and computes the C-index.
Usage
compute_prediction(
model,
test_data,
target_var = NULL,
trait.positive = NULL,
task_type = "classification",
time_var = NULL,
event_var = NULL,
file.name = NULL,
return = FALSE,
n_groups = 2
)Arguments
- model
The trained model returned as
$Modelbycompute_features.training.ML()(e.g.res$Model).- test_data
A data frame of predictor variables for the test set (for classification, a matrix is also accepted). It must contain all the features of the model; other columns are ignored for classification.
- target_var
Vector of true labels for the test set, in the order of the rows of
test_data(classification only).- trait.positive
Value in
target_varrepresenting the positive class (classification only).- task_type
Character. Either
"classification"or"survival".- time_var
Numeric vector of survival/follow-up times of the test samples, in the order of the rows of
test_data(required for survival tasks).- event_var
Numeric vector of event indicators of the test samples (1 = event, 0 = censored; required for survival tasks).
- file.name
Character. File name prefix of the plots saved in
Results/.- return
Logical. Whether to save the plots in
Results/(ROC and precision-recall curves for classification, Kaplan-Meier curves by predicted risk group for survival). Default = FALSE.- n_groups
Integer. Number of risk groups of the Kaplan-Meier plot (survival only, used with
return = TRUE). Default = 2. Seeplot_survival_performance().
Value
For classification, a list containing:
MetricsData frame of performance metrics (Accuracy, Sensitivity, Specificity, Precision, Recall, F1 score, MCC) for each threshold: one row per test sample, sorted by decreasing predicted probability.
AUCList with
AUROCandAUPRC, each a list with theestimateon the test set and thelowerandupperbounds of its 95% bootstrap confidence interval.PredictionsData frame of predicted probabilities for each class (columns
noandyes), in the order of the rows oftest_data.Curve_bandsList with
ROC(columnsfpr,lower,upper) andPRC(columnsrecall,lower,upper): pointwise 95% bootstrap bands around the ROC and precision-recall curves.
For survival, a list containing:
predsPredicted risk scores (higher values mean higher risk).
c_index,c_index_lower,c_index_upperC-index on the test set and its 95% confidence interval.
Details
Confidence intervals of AUROC and AUPRC come from 1000 bootstrap resamples of the test samples, stratified by class (positives and negatives resampled separately, so every resample has both), and the confidence interval of the C-index from 1000 bootstrap resamples stratified by event. Both use the fixed seed 123 only inside the bootstrap: the state of the random number generator of the session is not changed.
Survival models can predict a linear predictor, a survival time or a survival probability. The modelling library returns all three so that higher values mean longer survival (also the linear predictor of Cox models); they are reversed, so that higher predictions always mean higher risk.
Examples
if (FALSE) { # \dontrun{
data(data_example_classification)
X <- data_example_classification[, setdiff(colnames(data_example_classification), "target")]
y <- data_example_classification$target
set.seed(123)
train_idx <- caret::createDataPartition(y, p = 0.7, list = FALSE)
res <- compute_features.training.ML(features_train = X[train_idx, ],
target_var = y[train_idx],
task_type = "classification",
trait.positive = "1",
k_folds = 5,
n_rep = 2)
pred <- compute_prediction(model = res$Model,
test_data = X[-train_idx, ],
target_var = y[-train_idx],
task_type = "classification",
trait.positive = "1")
pred$AUC
} # }