Как работает решающее дерево для регрессии (Decision Tree Regressor)?

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

Ответ

Decision Tree Regressor — это модель, которая предсказывает непрерывную целевую переменную, рекурсивно разбивая пространство признаков на прямоугольные области (листья) и присваивая каждому листу постоянное значение — среднее значение целевой переменной обучающих выборок, попавших в этот лист.

Алгоритм построения (жадный):

  1. Начало: Весь набор данных считается корневым узлом.
  2. Выбор лучшего разделения: Для каждого признака и каждого возможного порогового значения (например, средние между соседними значениями) алгоритм вычисляет, насколько разделение уменьшит неоднородность в дочерних узлах.
    • Критерий для регрессии: Чаще всего используется дисперсия (или эквивалентно, среднеквадратичная ошибка (MSE) внутри узла). Алгоритм ищет разделение (признак, порог), которое максимизирует уменьшение дисперсии (или MSE).
    • Прирост информации (снижение MSE): MSE_parent - (weight_left * MSE_left + weight_right * MSE_right).
  3. Рекурсия: Шаг 2 рекурсивно повторяется для каждого дочернего узла, используя только данные, попавшие в этот узел.
  4. Критерии остановки: Рекурсия прекращается, когда выполняется одно из условий:
    • Достигнута max_depth (максимальная глубина).
    • В узле меньше min_samples_split образцов.
    • Снижение MSE меньше min_impurity_decrease.
    • В листе меньше min_samples_leaf образцов.
  5. Присвоение значения листу: Значением листа становится среднее арифметическое y всех обучающих образцов, попавших в этот лист.

Пример использования и визуализации в sklearn:

from sklearn.tree import DecisionTreeRegressor, plot_tree
import matplotlib.pyplot as plt

# Создание и обучение модели
regressor = DecisionTreeRegressor(max_depth=3, min_samples_leaf=5, random_state=42)
regressor.fit(X_train, y_train)

# Предсказание
y_pred = regressor.predict(X_test)

# Визуализация дерева (ключевое преимущество — интерпретируемость)
plt.figure(figsize=(12,8))
plot_tree(regressor, feature_names=X.columns, filled=True, rounded=True)
plt.show()

Сильные и слабые стороны:

  • Плюсы:

    • Высокая интерпретируемость. Можно проследить логику предсказания.
    • Не требует масштабирования/нормализации признаков.
    • Может моделировать нелинейные зависимости.
    • Работает с категориальными и числовыми признаками.
  • Минусы:

    • Высокая склонность к переобучению, особенно без ограничений глубины. Всегда требуется тщательная настройка гиперпараметров или использование в ансамблях (Random Forest, Gradient Boosting).
    • Нестабильность: Небольшие изменения в данных могут привести к построению совершенно другого дерева.
    • Плохая экстраполяция: Не может делать предсказания за пределами диапазона обучающих данных.