Перейти к содержимому

RandomForestClassifier

RandomForestClassifier — модель классификации, состоящая из многих деревьев решений.

Одно дерево решений строит правила и по ним выбирает класс. Случайный лес строит много разных деревьев, а затем объединяет их ответы голосованием.

RandomForestClassifier строит не одно дерево решений, а несколько деревьев.

Каждое дерево обучается немного по-разному и может дать свой вариант ответа. Затем ответы деревьев объединяются: для классификации выбирается класс, за который проголосовало большинство деревьев.

Такой подход называется ансамблем моделей. Ансамбль часто работает устойчивее, чем одно дерево, потому что итоговый ответ зависит не от одного набора правил, а от нескольких разных деревьев.

DecisionTreeClassifier хорошо объясняется, но одно дерево может переобучиться: слишком точно запомнить обучающие данные и хуже работать на новых объектах.

RandomForestClassifier обычно устойчивее, потому что отдельные деревья ошибаются по-разному. Голосование многих деревьев часто дает более надежный результат.

RandomForestClassifier хорошо подходит для задач классификации с большим числом признаков и сложными зависимостями.

Примеры:

  • распознавание рукописных цифр;
  • классификация объектов по множеству измерений;
  • табличные данные, где одно дерево дает нестабильный результат.

Используем датасет MnistSmall. В нем хранятся изображения рукописных цифр.

Каждый объект — это изображение цифры. Признаки — значения пикселей, а целевая переменная — правильная цифра от 0 до 9.

uses MLABC;
begin
var ds := Datasets.MnistSmall;
var df := ds.Data;
var X := df.ToMatrix(ds.Features);
var y := df.EncodeTarget(ds.Target);
var (Xtrain, Xtest, ytrain, ytest) :=
Validation.TrainTestSplit(X, y.Labels, testRatio := 0.2, seed := 42);
var model := new RandomForestClassifier(
nTrees := 20,
maxDepth := 20,
seed := 42
);
model.Fit(Xtrain, ytrain);
var pred := model.Predict(Xtest);
var acc := Metrics.Accuracy(ytest, pred);
Println('Accuracy:', acc:0:3);
end.

Вывод:

Accuracy: 0.910

Это означает, что модель правильно распознала примерно 91% цифр из тестовой выборки.

Для сравнения: на том же датасете одиночное дерево

var model := new DecisionTreeClassifier(
maxDepth := 20,
seed := 42
);

дает точность:

Accuracy: 0.732

А случайный лес из нескольких деревьев показывает заметно лучший результат.

  1. Загружается встроенный датасет MnistSmall.
  2. Признаки ds.Features преобразуются в матрицу X.
  3. Целевая переменная ds.Target кодируется в метки классов.
  4. Данные делятся на обучающую и тестовую выборки.
  5. Случайный лес обучается на Xtrain, ytrain.
  6. Качество проверяется на Xtest, ytest с помощью Accuracy.

nTrees — количество деревьев в лесу.

Чем больше деревьев, тем устойчивее может быть результат, но тем дольше обучение и предсказание.

maxDepth — максимальная глубина каждого дерева.

Если глубина слишком большая, деревья могут переобучаться. Если слишком маленькая, модель может быть слишком простой.

seed — начальное значение генератора случайных чисел.

Оно нужно, чтобы при каждом запуске получать одинаковый результат.

  • часто точнее одного дерева решений;
  • устойчивее к переобучению;
  • хорошо работает с большим числом признаков;
  • не требует сложной настройки для первого результата.
  • хуже объясняется, чем одно дерево;
  • обучается и работает медленнее одиночного дерева;
  • модель занимает больше памяти.

RandomForestClassifier полезен, когда одного дерева уже недостаточно. Он строит много деревьев и объединяет их ответы, поэтому часто дает более надежную классификацию.