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 чередуются.
Поэтому:
box) и «полосовых» графиках
(strip) значения признаков для классов A и B перекрываются
почти полностью, медианы близки;density) для двух классов похожи по
форме и положению, небольшие расхождения объясняются случайностью малой
выборки (50 наблюдений);pairs) и графике с эллипсами
(ellipse) точки двух классов перемешаны, эллипсы почти
совпадают, чёткой границы между классами нет и связи между признаками не
видно.Итог: ни один признак не разделяет классы, то есть на этих данных нет
информативных признаков. Это ожидаемо для случайного набора, и он
полезен как «нулевой» пример для сравнения с реальными данными
(например, 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, особенно
Sepal.Width) информативны заметно слабее;
Sepal.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
set.seed).Sepal.Length, тем чаще встречается setosa, а
высокие значения характерны для virginica; разные способы
дискретизации по-разному сохраняют это разделение.Набор 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"
shadowMin, shadowMean,
shadowMax) показывают важность случайных «теневых»
признаков, с которыми сравниваются реальные.caret доступно большое число моделей и инструментов
отбора признаков; графический анализ через featurePlot()
быстро показывает, различают ли признаки классы. На случайных данных
различий нет, и графики это наглядно подтверждают.FSelector позволяет оценить важность признаков
разными способами (фильтры, CFS, обёртки). На iris наиболее
информативны признаки лепестка, а Sepal.Width
наименее.arules зависит от метода:
interval сохраняет ширину интервалов,
frequency выравнивает число наблюдений,
cluster учитывает структуру данных, fixed даёт
полный контроль над границами.Ozone он
выделил группу признаков, связанных с температурой и сезоном.