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

GradientBoostingClassifier

GradientBoostingClassifier — ансамблевая модель классификации, построенная на деревьях решений.

Как и RandomForestClassifier, она использует не одно дерево, а много деревьев. Но деревья строятся по-другому.

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

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

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

Градиентный бустинг можно представить как последовательную работу над ошибками:

  1. первое дерево дает грубый прогноз;
  2. второе дерево исправляет ошибки первого;
  3. третье дерево исправляет ошибки суммы первых двух;
  4. следующие деревья продолжают уточнять модель.

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

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

Примеры:

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

Используем датасет MoscowHousing. Задача — предсказать, относится ли квартира к дорогим.

Для этого создадим новый столбец high_price: квартира считается дорогой, если ее цена выше медианной цены по датасету.

uses MLABC;
begin
var df := Datasets.MoscowHousing.Data;
var medianPrice := df.Median('price');
df := df.WithColumnBool(
'high_price',
row -> row['price'] > medianPrice
);
var features := [
'rooms',
'area',
'kitchen_area',
'floor',
'floors_total',
'metro_minutes'
];
var X := df.ToMatrix(features);
var target := df.EncodeTarget('high_price');
var (Xtrain, Xtest, ytrain, ytest) :=
Validation.TrainTestSplit(X, target.Labels, testRatio := 0.2, seed := 42);
var model := new GradientBoostingClassifier(
nEstimators := 50,
learningRate := 0.2,
maxDepth := 3,
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.880

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

  1. Загружается датасет MoscowHousing.
  2. Вычисляется медианная цена квартиры.
  3. Создается целевой столбец high_price.
  4. В признаки берутся только числовые столбцы.
  5. Данные делятся на обучающую и тестовую выборки.
  6. GradientBoostingClassifier обучается на тренировочной части.
  7. Качество проверяется с помощью Accuracy.

nEstimators — сколько деревьев построить.

Чем больше деревьев, тем дольше обучение. Иногда качество растет, но слишком большое число деревьев может привести к переобучению.

learningRate — насколько сильно учитывать каждое новое дерево.

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

maxDepth — насколько сложным может быть каждое дерево.

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

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

GradientBoostingClassifier строит цепочку деревьев. Каждое следующее дерево исправляет ошибки предыдущих, поэтому модель постепенно становится точнее.