На каком таргете (целевой переменной) учатся деревья в алгоритме бустинга для бинарной классификации?

«На каком таргете (целевой переменной) учатся деревья в алгоритме бустинга для бинарной классификации?» — вопрос из категории Деревья и ансамбли, который задают на 26% собеседований Data Scientist / ML Инженер. Ниже — развёрнутый ответ с разбором ключевых моментов.

Ответ

В алгоритмах градиентного бустинга (Gradient Boosting), таких как XGBoost, LightGBM или CatBoost, каждое новое дерево обучается не на исходной целевой переменной y, а на антиградиенте функции потерь (negative gradient), который часто называют псевдо-остатками.

Математическая суть: На шаге m ансамбль уже делает предсказание F_{m-1}(x). Мы хотим обучить новое дерево h_m(x), которое лучше всего приблизит разницу между истинным значением и текущим предсказанием. Эта разница вычисляется как градиент функции потерь L(y, F) по текущему предсказанию F_{m-1}(x).

Для бинарной классификации с лог-лоссом (logloss):

  • Текущее предсказание — это вероятность p_i = sigmoid(F_{m-1}(x_i)).
  • Псевдо-остаток (антиградиент) для объекта i вычисляется как: r_{im} = y_i - p_i где y_i ∈ {0, 1} — истинная метка класса.
  • Новое дерево h_m(x) обучается именно на этих остатках r_{im} как на целевых значениях.

Практический пример с XGBoost:

import xgboost as xgb
from sklearn.datasets import make_classification
from sklearn.model_selection import train_test_split

# Генерация данных
X, y = make_classification(n_samples=1000, n_features=20, n_classes=2)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2)

# Обучение модели
# Параметр 'objective' определяет функцию потерь, градиент которой вычисляется
model = xgb.XGBClassifier(
    objective='binary:logistic', # Используется лог-лосс
    n_estimators=100,
    learning_rate=0.1
)
model.fit(X_train, y_train)
# Каждое дерево внутри model обучалось на градиентах лог-лосса

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