Create trans_classifier object for machine-learning-based model prediction.
This class is a wrapper for methods of machine-learning-based classification models, including data pre-processing, feature selection, data split, model training, prediction, confusionMatrix and ROC (Receiver Operator Characteristic) or PR (Precision-Recall) curve.
Author(s): Felipe Mansoldo and Chi Liu
new()
Create the trans_classifier object.
trans_classifier$new( dataset = NULL, x.predictors = "all", y.response = NULL, n.cores = 1 )
datasetthe object of microtable Class.
x.predictorsdefault "all"; character string or data.frame; a character string represents selecting the corresponding data from microtable$taxa_abund; data.frame represents other customized data. See the following available options:
use all the taxa stored in microtable$taxa_abund
use Genus level table in microtable$taxa_abund, or other specific taxonomic rank, e.g. 'Phylum'
must be a data.frame; It should have the same format with the data.frame in microtable$taxa_abund, i.e. rows are features; cols are samples with same names in sample_table
y.responsedefault NULL; the response variable in sample_table.
n.coresdefault 1; the CPU thread used.
data_feature and data_response in the object.
\donttest{
data(dataset)
t1 <- trans_classifier$new(
dataset = dataset,
x.predictors = "Genus",
y.response = "Group")
}
cal_preProcess()
Pre-process (centering, scaling etc.) of the feature data based on the caret::preProcess function. See https://topepo.github.io/caret/pre-processing.html for more details.
trans_classifier$cal_preProcess(...)
...parameters pass to preProcess function of caret package.
converted data_feature in the object.
\dontrun{
t1$cal_preProcess(method = c("center", "scale", "nzv"))
}
cal_feature_sel()
Perform feature selection. See https://topepo.github.io/caret/feature-selection-overview.html for more details.
trans_classifier$cal_feature_sel( boruta.maxRuns = 300, boruta.pValue = 0.01, boruta.repetitions = 4, ... )
boruta.maxRunsdefault 300; maximal number of importance source runs; passed to the maxRuns parameter in Boruta function of Boruta package.
boruta.pValuedefault 0.01; p value passed to the pValue parameter in Boruta function of Boruta package.
boruta.repetitionsdefault 4; repetition runs for the feature selection.
...parameters pass to Boruta function of Boruta package.
optimized data_feature in the object.
\donttest{
t1$cal_feature_sel(boruta.maxRuns = 300, boruta.pValue = 0.01)
}
cal_split()
Split data for training and testing.
trans_classifier$cal_split(prop.train = 3/4)
prop.traindefault 3/4; the ratio of the dataset used for the training.
data_train and data_test in the object.
\donttest{
t1$cal_split(prop.train = 3/4)
}
set_trainControl()
Control parameters for the following training. See trainControl function of caret package for details.
trans_classifier$set_trainControl( method = "repeatedcv", classProbs = TRUE, savePredictions = TRUE, ... )
methoddefault 'repeatedcv'; 'repeatedcv': Repeated k-Fold cross validation; see method parameter in trainControl function of caret package for available options.
classProbsdefault TRUE; should class probabilities be computed for classification models?; see classProbs parameter in caret::trainControl function.
savePredictionsdefault TRUE; see savePredictions parameter in caret::trainControl function
...parameters pass to trainControl function of caret package.
trainControl in the object.
\dontrun{
t1$set_trainControl(method = 'repeatedcv')
}
cal_train()
Run the model training.
trans_classifier$cal_train( method = "rf", metric = "Accuracy", max.mtry = 2, max.ntree = 200, ... )
methoddefault "rf"; "rf": random forest; see method in caret::train function for other options.
metricdefault "Accuracy"; see metric in caret::train function for other options.
max.mtrydefault 2; for method = "rf"; maximum mtry used for the tunegrid to do hyperparameter tuning to optimize the model.
max.ntreedefault 200; for method = "rf"; maximum number of trees used to optimize the model.
...parameters pass to train function of caret package.
res_train in the object.
\dontrun{
# random forest
t1$cal_train(method = "rf")
# Support Vector Machines with Radial Basis Function Kernel
t1$cal_train(method = "svmRadial", tuneLength = 15)
}
cal_feature_imp()
Get feature importance from the training model.
trans_classifier$cal_feature_imp(...)
...parameters pass to varImp function of caret package.
res_feature_imp in the object. One row for each predictor variable. The column(s) are different importance measures.
\dontrun{
t1$cal_feature_imp()
}
cal_predict()
Run the prediction.
trans_classifier$cal_predict(positive_class = NULL)
positive_classdefault NULL; see positive parameter in confusionMatrix function of caret package; If positive_class is NULL, use the first group in data as the positive class automatically.
res_predict, res_confusion_fit and res_confusion_stats stored in the object.
\dontrun{
t1$cal_predict()
}
plot_confusionMatrix()
Plot the cross-tabulation of observed and predicted classes with associated statistics.
trans_classifier$plot_confusionMatrix( plot_confusion = TRUE, plot_statistics = TRUE )
plot_confusiondefault TRUE; whether plot the confusion matrix.
plot_statisticsdefault TRUE; whether plot the statistics.
ggplot object.
\dontrun{
t1$plot_confusionMatrix()
}
cal_ROC()
Get ROC (Receiver Operator Characteristic) curve data and the performance data.
trans_classifier$cal_ROC(input = "pred")
inputdefault "pred"; 'pred' or 'train'; 'pred' represents using prediction results; 'train' represents using training results.
a list res_ROC stored in the object.
\dontrun{
t1$cal_ROC()
}
plot_ROC()
Plot ROC curve.
trans_classifier$plot_ROC(
plot_type = c("ROC", "PR")[1],
plot_group = "all",
color_values = RColorBrewer::brewer.pal(8, "Dark2"),
add_AUC = TRUE,
plot_method = FALSE,
...
)plot_typedefault c("ROC", "PR")[1]; 'ROC' represents ROC (Receiver Operator Characteristic) curve; 'PR' represents PR (Precision-Recall) curve.
plot_groupdefault "all"; 'all' represents all the classes in the model; 'add' represents all adding micro-average and macro-average results, see https://scikit-learn.org/stable/auto_examples/model_selection/plot_roc.html; other options should be one or more class names, same with the names in Group column of res_ROC$res_roc from cal_ROC function.
color_valuesdefault RColorBrewer::brewer.pal(8, "Dark2"); colors used in the plot.
add_AUCdefault TRUE; whether add AUC in the legend.
plot_methoddefault FALSE; If TRUE, show the method in the legend though only one method is found.
...parameters pass to geom_path function of ggplot2 package.
ggplot2 object.
\dontrun{
t1$plot_ROC(size = 1, alpha = 0.7)
}
clone()
The objects of this class are cloneable with this method.
trans_classifier$clone(deep = FALSE)
deepWhether to make a deep clone.
## ------------------------------------------------
## Method `trans_classifier$new`
## ------------------------------------------------
data(dataset)
t1 <- trans_classifier$new(
dataset = dataset,
x.predictors = "Genus",
y.response = "Group")
## ------------------------------------------------
## Method `trans_classifier$cal_preProcess`
## ------------------------------------------------
## Not run:
t1$cal_preProcess(method = c("center", "scale", "nzv"))
## End(Not run)
## ------------------------------------------------
## Method `trans_classifier$cal_feature_sel`
## ------------------------------------------------
t1$cal_feature_sel(boruta.maxRuns = 300, boruta.pValue = 0.01)
## ------------------------------------------------
## Method `trans_classifier$cal_split`
## ------------------------------------------------
t1$cal_split(prop.train = 3/4)
## ------------------------------------------------
## Method `trans_classifier$set_trainControl`
## ------------------------------------------------
## Not run:
t1$set_trainControl(method = 'repeatedcv')
## End(Not run)
## ------------------------------------------------
## Method `trans_classifier$cal_train`
## ------------------------------------------------
## Not run:
# random forest
t1$cal_train(method = "rf")
# Support Vector Machines with Radial Basis Function Kernel
t1$cal_train(method = "svmRadial", tuneLength = 15)
## End(Not run)
## ------------------------------------------------
## Method `trans_classifier$cal_feature_imp`
## ------------------------------------------------
## Not run:
t1$cal_feature_imp()
## End(Not run)
## ------------------------------------------------
## Method `trans_classifier$cal_predict`
## ------------------------------------------------
## Not run:
t1$cal_predict()
## End(Not run)
## ------------------------------------------------
## Method `trans_classifier$plot_confusionMatrix`
## ------------------------------------------------
## Not run:
t1$plot_confusionMatrix()
## End(Not run)
## ------------------------------------------------
## Method `trans_classifier$cal_ROC`
## ------------------------------------------------
## Not run:
t1$cal_ROC()
## End(Not run)
## ------------------------------------------------
## Method `trans_classifier$plot_ROC`
## ------------------------------------------------
## Not run:
t1$plot_ROC(size = 1, alpha = 0.7)
## End(Not run)Please choose more modern alternatives, such as Google Chrome or Mozilla Firefox.