library(caret)
names(getModelInfo())
##   [1] "ada"                 "AdaBag"              "AdaBoost.M1"        
##   [4] "adaboost"            "amdai"               "ANFIS"              
##   [7] "avNNet"              "awnb"                "awtan"              
##  [10] "bag"                 "bagEarth"            "bagEarthGCV"        
##  [13] "bagFDA"              "bagFDAGCV"           "bam"                
##  [16] "bartMachine"         "bayesglm"            "binda"              
##  [19] "blackboost"          "blasso"              "blassoAveraged"     
##  [22] "bridge"              "brnn"                "BstLm"              
##  [25] "bstSm"               "bstTree"             "C5.0"               
##  [28] "C5.0Cost"            "C5.0Rules"           "C5.0Tree"           
##  [31] "cforest"             "chaid"               "CSimca"             
##  [34] "ctree"               "ctree2"              "cubist"             
##  [37] "dda"                 "deepboost"           "DENFIS"             
##  [40] "dnn"                 "dwdLinear"           "dwdPoly"            
##  [43] "dwdRadial"           "earth"               "elm"                
##  [46] "enet"                "evtree"              "extraTrees"         
##  [49] "fda"                 "FH.GBML"             "FIR.DM"             
##  [52] "foba"                "FRBCS.CHI"           "FRBCS.W"            
##  [55] "FS.HGD"              "gam"                 "gamboost"           
##  [58] "gamLoess"            "gamSpline"           "gaussprLinear"      
##  [61] "gaussprPoly"         "gaussprRadial"       "gbm_h2o"            
##  [64] "gbm"                 "gcvEarth"            "GFS.FR.MOGUL"       
##  [67] "GFS.LT.RS"           "GFS.THRIFT"          "glm.nb"             
##  [70] "glm"                 "glmboost"            "glmnet_h2o"         
##  [73] "glmnet"              "glmStepAIC"          "gpls"               
##  [76] "hda"                 "hdda"                "hdrda"              
##  [79] "HYFIS"               "icr"                 "J48"                
##  [82] "JRip"                "kernelpls"           "kknn"               
##  [85] "knn"                 "krlsPoly"            "krlsRadial"         
##  [88] "lars"                "lars2"               "lasso"              
##  [91] "lda"                 "lda2"                "leapBackward"       
##  [94] "leapForward"         "leapSeq"             "Linda"              
##  [97] "lm"                  "lmStepAIC"           "LMT"                
## [100] "loclda"              "logicBag"            "LogitBoost"         
## [103] "logreg"              "lssvmLinear"         "lssvmPoly"          
## [106] "lssvmRadial"         "lvq"                 "M5"                 
## [109] "M5Rules"             "manb"                "mda"                
## [112] "Mlda"                "mlp"                 "mlpKerasDecay"      
## [115] "mlpKerasDecayCost"   "mlpKerasDropout"     "mlpKerasDropoutCost"
## [118] "mlpML"               "mlpSGD"              "mlpWeightDecay"     
## [121] "mlpWeightDecayML"    "monmlp"              "msaenet"            
## [124] "multinom"            "mxnet"               "mxnetAdam"          
## [127] "naive_bayes"         "nb"                  "nbDiscrete"         
## [130] "nbSearch"            "neuralnet"           "nnet"               
## [133] "nnls"                "nodeHarvest"         "null"               
## [136] "OneR"                "ordinalNet"          "ordinalRF"          
## [139] "ORFlog"              "ORFpls"              "ORFridge"           
## [142] "ORFsvm"              "ownn"                "pam"                
## [145] "parRF"               "PART"                "partDSA"            
## [148] "pcaNNet"             "pcr"                 "pda"                
## [151] "pda2"                "penalized"           "PenalizedLDA"       
## [154] "plr"                 "pls"                 "plsRglm"            
## [157] "polr"                "ppr"                 "pre"                
## [160] "PRIM"                "protoclass"          "qda"                
## [163] "QdaCov"              "qrf"                 "qrnn"               
## [166] "randomGLM"           "ranger"              "rbf"                
## [169] "rbfDDA"              "Rborist"             "rda"                
## [172] "regLogistic"         "relaxo"              "rf"                 
## [175] "rFerns"              "RFlda"               "rfRules"            
## [178] "ridge"               "rlda"                "rlm"                
## [181] "rmda"                "rocc"                "rotationForest"     
## [184] "rotationForestCp"    "rpart"               "rpart1SE"           
## [187] "rpart2"              "rpartCost"           "rpartScore"         
## [190] "rqlasso"             "rqnc"                "RRF"                
## [193] "RRFglobal"           "rrlda"               "RSimca"             
## [196] "rvmLinear"           "rvmPoly"             "rvmRadial"          
## [199] "SBC"                 "sda"                 "sdwd"               
## [202] "simpls"              "SLAVE"               "slda"               
## [205] "smda"                "snn"                 "sparseLDA"          
## [208] "spikeslab"           "spls"                "stepLDA"            
## [211] "stepQDA"             "superpc"             "svmBoundrangeString"
## [214] "svmExpoString"       "svmLinear"           "svmLinear2"         
## [217] "svmLinear3"          "svmLinearWeights"    "svmLinearWeights2"  
## [220] "svmPoly"             "svmRadial"           "svmRadialCost"      
## [223] "svmRadialSigma"      "svmRadialWeights"    "svmSpectrumString"  
## [226] "tan"                 "tanSearch"           "treebag"            
## [229] "vbmpRadial"          "vglmAdjCat"          "vglmContRatio"      
## [232] "vglmCumulative"      "widekernelpls"       "WM"                 
## [235] "wsrf"                "xgbDART"             "xgbLinear"          
## [238] "xgbTree"             "xyf"

Всего в пакете caret доступно 239 моделей. Среди них можно выделить те, у которых в описании (поле tags) указан отбор признаков (встроенный, как у деревьев, лассо, random forest, или неявный):

models <- getModelInfo()
fs_models <- names(models)[sapply(models, function(m)
  any(grepl("Feature Selection", m$tags)))]
length(fs_models)
## [1] 90
fs_models
##  [1] "ada"              "AdaBag"           "AdaBoost.M1"      "adaboost"        
##  [5] "bagEarth"         "bagEarthGCV"      "bagFDA"           "bagFDAGCV"       
##  [9] "bartMachine"      "blasso"           "BstLm"            "bstSm"           
## [13] "C5.0"             "C5.0Cost"         "C5.0Rules"        "C5.0Tree"        
## [17] "cforest"          "chaid"            "ctree"            "ctree2"          
## [21] "cubist"           "deepboost"        "earth"            "enet"            
## [25] "evtree"           "extraTrees"       "fda"              "foba"            
## [29] "gamboost"         "gbm_h2o"          "gbm"              "gcvEarth"        
## [33] "glmnet_h2o"       "glmnet"           "glmStepAIC"       "J48"             
## [37] "JRip"             "lars"             "lars2"            "lasso"           
## [41] "leapBackward"     "leapForward"      "leapSeq"          "lmStepAIC"       
## [45] "LMT"              "LogitBoost"       "M5"               "M5Rules"         
## [49] "msaenet"          "nodeHarvest"      "OneR"             "ordinalNet"      
## [53] "ordinalRF"        "ORFlog"           "ORFpls"           "ORFridge"        
## [57] "ORFsvm"           "pam"              "parRF"            "PART"            
## [61] "penalized"        "PenalizedLDA"     "qrf"              "ranger"          
## [65] "Rborist"          "relaxo"           "rf"               "rFerns"          
## [69] "rfRules"          "rotationForest"   "rotationForestCp" "rpart"           
## [73] "rpart1SE"         "rpart2"           "rpartCost"        "rpartScore"      
## [77] "rqlasso"          "rqnc"             "RRF"              "RRFglobal"       
## [81] "sdwd"             "smda"             "sparseLDA"        "spikeslab"       
## [85] "stepLDA"          "stepQDA"          "wsrf"             "xgbDART"         
## [89] "xgbLinear"        "xgbTree"

Кроме того, в caret есть отдельные функции для отбора признаков: rfe() (рекурсивное исключение признаков), sbf() (фильтрация по одномерным статистикам), gafs() (генетический алгоритм) и safs() (имитация отжига), а также findCorrelation(), nearZeroVar() и varImp().

set.seed(1)
x <- matrix(rnorm(50 * 5), ncol = 5)
colnames(x) <- paste0("X", 1:5)
x <- as.data.frame(x)
y <- factor(rep(c("A", "B"), 25))

str(x)
## 'data.frame':    50 obs. of  5 variables:
##  $ X1: num  -0.626 0.184 -0.836 1.595 0.33 ...
##  $ X2: num  0.398 -0.612 0.341 -1.129 1.433 ...
##  $ X3: num  -0.6204 0.0421 -0.9109 0.158 -0.6546 ...
##  $ X4: num  0.4502 -0.0186 -0.3181 -0.9294 -1.4875 ...
##  $ X5: num  0.409 1.689 1.587 -0.331 -2.285 ...
table(y)
## y
##  A  B 
## 25 25
dir.create("plots", showWarnings = FALSE)

plot_types <- c("strip", "box", "density", "pairs", "ellipse")

for (tp in plot_types) {
  p <- switch(tp,
    strip   = featurePlot(x, y, plot = "strip", jitter = TRUE,
                          auto.key = list(columns = 2)),
    box     = featurePlot(x, y, plot = "box",
                          scales = list(y = list(relation = "free"),
                                        x = list(rot = 90)),
                          layout = c(5, 1), auto.key = list(columns = 2)),
    density = featurePlot(x, y, plot = "density",
                          scales = list(x = list(relation = "free"),
                                        y = list(relation = "free")),
                          adjust = 1.5, pch = "|", layout = c(5, 1),
                          auto.key = list(columns = 2)),
    pairs   = featurePlot(x, y, plot = "pairs",
                          auto.key = list(columns = 2)),
    ellipse = featurePlot(x, y, plot = "ellipse",
                          auto.key = list(columns = 2))
  )
  jpeg(file.path("plots", paste0("featurePlot_", tp, ".jpg")),
       width = 1000, height = 700, res = 110)
  print(p)
  dev.off()
}

list.files("plots", pattern = "\\.jpg$")
## [1] "featurePlot_box.jpg"     "featurePlot_density.jpg"
## [3] "featurePlot_ellipse.jpg" "featurePlot_pairs.jpg"  
## [5] "featurePlot_strip.jpg"

Сохранённые изображения (те же файлы, что лежат в папке plots):

strip

box

density

pairs

ellipse

Данные сгенерированы случайно: признаки X1–X5 взяты из стандартного нормального распределения независимо от класса, а метки A и B чередуются. Поэтому:

Итог: ни один признак не разделяет классы, то есть на этих данных нет информативных признаков. Это ожидаемо для случайного набора, и он полезен как «нулевой» пример для сравнения с реальными данными (например, iris), где графики выглядят иначе.

library(FSelector)
data(iris)
str(iris)
## 'data.frame':    150 obs. of  5 variables:
##  $ Sepal.Length: num  5.1 4.9 4.7 4.6 5 5.4 4.6 5 4.4 4.9 ...
##  $ Sepal.Width : num  3.5 3 3.2 3.1 3.6 3.9 3.4 3.4 2.9 3.1 ...
##  $ Petal.Length: num  1.4 1.4 1.3 1.5 1.4 1.7 1.4 1.5 1.4 1.5 ...
##  $ Petal.Width : num  0.2 0.2 0.2 0.2 0.2 0.4 0.3 0.2 0.2 0.1 ...
##  $ Species     : Factor w/ 3 levels "setosa","versicolor",..: 1 1 1 1 1 1 1 1 1 1 ...
set.seed(123)

fs_list <- list(
  chi.squared = chi.squared(Species ~ ., iris),
  information.gain = information.gain(Species ~ ., iris),
  gain.ratio = gain.ratio(Species ~ ., iris),
  symmetrical.uncertainty = symmetrical.uncertainty(Species ~ ., iris),
  oneR = oneR(Species ~ ., iris),
  relief = relief(Species ~ ., iris, neighbours.count = 5, sample.size = 20),
  random.forest = random.forest.importance(Species ~ ., iris)
)

imp <- sapply(fs_list, function(d) d$attr_importance)
rownames(imp) <- rownames(fs_list[[1]])
round(imp, 4)
##              chi.squared information.gain gain.ratio symmetrical.uncertainty
## Sepal.Length      0.6288           0.4521     0.4196                  0.4156
## Sepal.Width       0.4922           0.2673     0.2473                  0.2453
## Petal.Length      0.9346           0.9403     0.8585                  0.8572
## Petal.Width       0.9432           0.9554     0.8714                  0.8705
##                oneR relief random.forest
## Sepal.Length 0.1733 0.1871       14.7858
## Sepal.Width  0.0400 0.1142        5.8739
## Petal.Length 0.4000 0.3645       49.6865
## Petal.Width  0.4067 0.3700       46.6599

Чтобы сравнить методы, приведём важность к шкале от 0 до 1 (деление на максимум по методу):

imp_norm <- apply(imp, 2, function(z) z / max(z))
round(imp_norm, 3)
##              chi.squared information.gain gain.ratio symmetrical.uncertainty
## Sepal.Length       0.667            0.473      0.482                   0.477
## Sepal.Width        0.522            0.280      0.284                   0.282
## Petal.Length       0.991            0.984      0.985                   0.985
## Petal.Width        1.000            1.000      1.000                   1.000
##               oneR relief random.forest
## Sepal.Length 0.426  0.506         0.298
## Sepal.Width  0.098  0.309         0.118
## Petal.Length 0.984  0.985         1.000
## Petal.Width  1.000  1.000         0.939
barplot(t(imp_norm), beside = TRUE, las = 2, ylim = c(0, 1.5),
        col = rainbow(ncol(imp_norm)),
        legend.text = colnames(imp_norm),
        args.legend = list(x = "topleft", cex = 0.6, ncol = 2),
        main = "Нормированная важность признаков iris")

Средний ранг признака по всем методам (1 = самый важный):

ranks <- apply(-imp_norm, 2, rank, ties.method = "average")
sort(rowMeans(ranks))
##  Petal.Width Petal.Length Sepal.Length  Sepal.Width 
##     1.142857     1.857143     3.000000     4.000000

Функция cutoff.k() берёт k лучших признаков по выбранной оценке:

w <- information.gain(Species ~ ., iris)
best2 <- cutoff.k(w, 2)
best2
## [1] "Petal.Width"  "Petal.Length"
as.simple.formula(best2, "Species")
## Species ~ Petal.Width + Petal.Length
## <environment: 0x000001e41f730120>

Метод CFS (correlation-based feature selection) сразу возвращает подмножество:

cfs_subset <- cfs(Species ~ ., iris)
cfs_subset
## [1] "Petal.Length" "Petal.Width"

Метод-обёртка: жадный поиск вперёд с оценкой качества подмножества деревом решений (5-кратная кросс-валидация):

library(rpart)
set.seed(123)

evaluator <- function(subset) {
  k <- 5
  splits <- runif(nrow(iris))
  results <- sapply(1:k, function(i) {
    test_idx <- (splits >= (i - 1) / k) & (splits < i / k)
    train <- iris[!test_idx, , drop = FALSE]
    test  <- iris[test_idx, , drop = FALSE]
    tree <- rpart(as.simple.formula(subset, "Species"), train)
    err <- mean(test$Species != predict(tree, test, type = "class"))
    1 - err
  })
  mean(results)
}

fwd <- forward.search(names(iris)[-5], evaluator)
fwd
## [1] "Sepal.Length" "Petal.Width"

Возьмём переменную Sepal.Length набора iris.

library(arules)
set.seed(123)

v <- iris$Sepal.Length
summary(v)
##    Min. 1st Qu.  Median    Mean 3rd Qu.    Max. 
##   4.300   5.100   5.800   5.843   6.400   7.900
d_interval  <- discretize(v, method = "interval",  breaks = 3)
d_frequency <- discretize(v, method = "frequency", breaks = 3)
d_cluster   <- discretize(v, method = "cluster",   breaks = 3)
d_fixed     <- discretize(v, method = "fixed",
                          breaks = c(-Inf, 5.5, 6.5, Inf))

lapply(list(interval = d_interval, frequency = d_frequency,
            cluster = d_cluster, fixed = d_fixed), table)
## $interval
## 
## [4.3,5.5) [5.5,6.7) [6.7,7.9] 
##        52        70        28 
## 
## $frequency
## 
## [4.3,5.4) [5.4,6.3) [6.3,7.9] 
##        46        53        51 
## 
## $cluster
## 
##  [4.3,5.37) [5.37,6.36)  [6.36,7.9] 
##          46          62          42 
## 
## $fixed
## 
## [-Inf,5.5)  [5.5,6.5) [6.5, Inf] 
##         52         63         35
cuts <- list(
  interval  = discretize(v, method = "interval",  breaks = 3, onlycuts = TRUE),
  frequency = discretize(v, method = "frequency", breaks = 3, onlycuts = TRUE),
  cluster   = discretize(v, method = "cluster",   breaks = 3, onlycuts = TRUE),
  fixed     = c(-Inf, 5.5, 6.5, Inf)
)
cuts
## $interval
## [1] 4.3 5.5 6.7 7.9
## 
## $frequency
## [1] 4.3 5.4 6.3 7.9
## 
## $cluster
## [1] 4.300000 5.332732 6.272161 7.900000
## 
## $fixed
## [1] -Inf  5.5  6.5  Inf
op <- par(mfrow = c(2, 2))
for (m in names(cuts)) {
  hist(v, breaks = 20, col = "grey85", border = "white",
       main = paste("Метод:", m), xlab = "Sepal.Length")
  abline(v = cuts[[m]][is.finite(cuts[[m]])], col = "red", lwd = 2, lty = 2)
}

par(op)
lapply(list(interval = d_interval, frequency = d_frequency,
            cluster = d_cluster, fixed = d_fixed),
       function(d) table(d, iris$Species))
## $interval
##            
## d           setosa versicolor virginica
##   [4.3,5.5)     45          6         1
##   [5.5,6.7)      5         38        27
##   [6.7,7.9]      0          6        22
## 
## $frequency
##            
## d           setosa versicolor virginica
##   [4.3,5.4)     40          5         1
##   [5.4,6.3)     10         31        12
##   [6.3,7.9]      0         14        37
## 
## $cluster
##              
## d             setosa versicolor virginica
##   [4.3,5.37)      40          5         1
##   [5.37,6.36)     10         34        18
##   [6.36,7.9]       0         11        31
## 
## $fixed
##             
## d            setosa versicolor virginica
##   [-Inf,5.5)     45          6         1
##   [5.5,6.5)       5         35        23
##   [6.5, Inf]      0          9        26

Набор Ozone из пакета mlbench содержит ежедневные измерения озона в Лос-Анджелесе (1976 г.). Целевая переменная V4, дневной максимум почасового среднего значения озона. Остальные переменные:

Переменная Смысл
V1 месяц
V2 день месяца
V3 день недели
V5 высота изобарической поверхности 500 мбар (Vandenberg AFB)
V6 скорость ветра (LAX)
V7 влажность (LAX)
V8 температура (Sandburg)
V9 температура (El Monte)
V10 высота нижней границы инверсии (LAX)
V11 градиент давления (LAX–Daggett)
V12 температура на границе инверсии (LAX)
V13 видимость (LAX)
library(mlbench)
library(Boruta)

data("Ozone", package = "mlbench")
dim(Ozone)
## [1] 366  13
colSums(is.na(Ozone))
##  V1  V2  V3  V4  V5  V6  V7  V8  V9 V10 V11 V12 V13 
##   0   0   0   5  12   0  15   2 139  15   1  14   0
Ozone_c <- na.omit(Ozone)
dim(Ozone_c)
## [1] 203  13

Boruta не работает с пропусками, поэтому строки с NA удалены: осталось 203 наблюдений из 366.

set.seed(123)
b <- Boruta(V4 ~ ., data = Ozone_c, doTrace = 0)
print(b)
## Boruta performed 18 iterations in 0.362658 secs.
##  9 attributes confirmed important: V1, V10, V11, V12, V13 and 4 more;
##  3 attributes confirmed unimportant: V2, V3, V6;
attStats(b)
plot(b, xlab = "", xaxt = "n", main = "Boruta: важность признаков (Ozone)")
lz <- lapply(1:ncol(b$ImpHistory), function(i)
  b$ImpHistory[is.finite(b$ImpHistory[, i]), i])
names(lz) <- colnames(b$ImpHistory)
Labels <- sort(sapply(lz, median))
axis(side = 1, las = 2, labels = names(Labels),
     at = 1:ncol(b$ImpHistory), cex.axis = 0.7)

Динамика важности по итерациям:

plotImpHistory(b)

Если остались неопределённые («tentative») признаки, их можно классифицировать дополнительно:

b_fix <- TentativeRoughFix(b)
print(b_fix)
## Boruta performed 18 iterations in 0.362658 secs.
##  9 attributes confirmed important: V1, V10, V11, V12, V13 and 4 more;
##  3 attributes confirmed unimportant: V2, V3, V6;
getSelectedAttributes(b_fix, withTentative = FALSE)
## [1] "V1"  "V5"  "V7"  "V8"  "V9"  "V10" "V11" "V12" "V13"
  1. В caret доступно большое число моделей и инструментов отбора признаков; графический анализ через featurePlot() быстро показывает, различают ли признаки классы. На случайных данных различий нет, и графики это наглядно подтверждают.
  2. Пакет FSelector позволяет оценить важность признаков разными способами (фильтры, CFS, обёртки). На iris наиболее информативны признаки лепестка, а Sepal.Width наименее.
  3. Дискретизация в arules зависит от метода: interval сохраняет ширину интервалов, frequency выравнивает число наблюдений, cluster учитывает структуру данных, fixed даёт полный контроль над границами.
  4. Алгоритм Boruta сравнивает реальные признаки с «теневыми» копиями и надёжно отделяет значимые признаки от шума; на Ozone он выделил группу признаков, связанных с температурой и сезоном.