Spec-Zone.ru › scikit-learn

Примечание

Перейти к концу, чтобы загрузить весь пример кода. или запустить этот пример в браузере через JupyterLite или Binder

Преждевременная остановка в градиентном бустинге

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

Преждевременная остановка — это техника в градиентном бустинге, которая позволяет нам найти оптимальное количество итераций, необходимое для построения модели, которая хорошо обобщается на невидимые данные и избегает переобучения. Идея проста: мы откладываем часть наших данных как валидационный набор (указанный с помощью validation_fraction) для оценки производительности модели во время обучения. По мере итеративного построения модели с добавлением новых этапов (деревьев), её производительность на валидационном наборе отслеживается как функция от количества шагов.

Преждевременная остановка становится эффективной, когда производительность модели на валидационном наборе достигает плато или ухудшается (в пределах отклонений, указанных tol) на определённом количестве последовательных этапов (указанном n_iter_no_change). Это сигнализирует о том, что модель достигла точки, где дальнейшие итерации могут привести к переобучению, и пора остановить обучение.

Количество оценщиков (деревьев) в окончательной модели, когда применяется преждевременная остановка, можно получить, используя атрибут n_estimators_. В целом, преждевременная остановка является ценным инструментом для достижения баланса между производительностью модели и эффективностью в градиентном бустинге.

Лицензия: BSD 3-х пунктов

# Authors: The scikit-learn developers
# SPDX-License-Identifier: BSD-3-Clause

Подготовка данных

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

import time

import matplotlib.pyplot as plt

from sklearn.datasets import fetch_california_housing
from sklearn.ensemble import GradientBoostingRegressor
from sklearn.metrics import mean_squared_error
from sklearn.model_selection import train_test_split

data = fetch_california_housing()
X, y = data.data[:600], data.target[:600]

X_train, X_val, y_train, y_val = train_test_split(X, y, test_size=0.2, random_state=42)

Обучение и сравнение моделей

Обучаются две модели GradientBoostingRegressor: одна с преждевременной остановкой, а другая — без неё. Цель — сравнить их производительность. Также рассчитывается время обучения и n_estimators_, используемые обеими моделями.

params = dict(n_estimators=1000, max_depth=5, learning_rate=0.1, random_state=42)

gbm_full = GradientBoostingRegressor(**params)
gbm_early_stopping = GradientBoostingRegressor(
    **params,
    validation_fraction=0.1,
    n_iter_no_change=10,
)

start_time = time.time()
gbm_full.fit(X_train, y_train)
training_time_full = time.time() - start_time
n_estimators_full = gbm_full.n_estimators_

start_time = time.time()
gbm_early_stopping.fit(X_train, y_train)
training_time_early_stopping = time.time() - start_time
estimators_early_stopping = gbm_early_stopping.n_estimators_

Вычисление ошибок

Код вычисляет mean_squared_error для обучающего и валидационного наборов данных для моделей, обученных в предыдущем разделе. Вычисляются ошибки для каждой итерации бустинга. Цель — оценить производительность и сходимость моделей.

train_errors_without = []
val_errors_without = []

train_errors_with = []
val_errors_with = []

for i, (train_pred, val_pred) in enumerate(
    zip(
        gbm_full.staged_predict(X_train),
        gbm_full.staged_predict(X_val),
    )
):
    train_errors_without.append(mean_squared_error(y_train, train_pred))
    val_errors_without.append(mean_squared_error(y_val, val_pred))

for i, (train_pred, val_pred) in enumerate(
    zip(
        gbm_early_stopping.staged_predict(X_train),
        gbm_early_stopping.staged_predict(X_val),
    )
):
    train_errors_with.append(mean_squared_error(y_train, train_pred))
    val_errors_with.append(mean_squared_error(y_val, val_pred))

Визуализация сравнения

Включает три подграфика:

  1. График ошибок обучения обеих моделей по итерациям бустинга.
  2. График ошибок валидации обеих моделей по итерациям бустинга.
  3. Столбчатая диаграмма для сравнения времен обучения и количества оценщиков моделей с и без преждевременной остановки.
fig, axes = plt.subplots(ncols=3, figsize=(12, 4))

axes[0].plot(train_errors_without, label="gbm_full")
axes[0].plot(train_errors_with, label="gbm_early_stopping")
axes[0].set_xlabel("Boosting Iterations")
axes[0].set_ylabel("MSE (Training)")
axes[0].set_yscale("log")
axes[0].legend()
axes[0].set_title("Training Error")

axes[1].plot(val_errors_without, label="gbm_full")
axes[1].plot(val_errors_with, label="gbm_early_stopping")
axes[1].set_xlabel("Boosting Iterations")
axes[1].set_ylabel("MSE (Validation)")
axes[1].set_yscale("log")
axes[1].legend()
axes[1].set_title("Validation Error")

training_times = [training_time_full, training_time_early_stopping]
labels = ["gbm_full", "gbm_early_stopping"]
bars = axes[2].bar(labels, training_times)
axes[2].set_ylabel("Training Time (s)")

for bar, n_estimators in zip(bars, [n_estimators_full, estimators_early_stopping]):
    height = bar.get_height()
    axes[2].text(
        bar.get_x() + bar.get_width() / 2,
        height + 0.001,
        f"Estimators: {n_estimators}",
        ha="center",
        va="bottom",
    )

plt.tight_layout()
plt.show()
Training Error, Validation Error

Разница в ошибке обучения между gbm_full и gbm_early_stopping обусловлена тем, что gbm_early_stopping выделяет validation_fraction обучающих данных для внутренней валидации. Преждевременная остановка определяется на основе этого внутреннего валидационного результата.

Итоги

В нашем примере с моделью GradientBoostingRegressor на наборе данных о ценах на жильё в Калифорнии мы продемонстрировали практические преимущества преждевременной остановки:

  • Предотвращение переобучения: Мы показали, как ошибка валидации стабилизируется или начинает увеличиваться после определённого момента, что свидетельствует о лучшем обобщении модели на невидимые данные. Это достигается остановкой процесса обучения до возникновения переобучения.
  • Повышение эффективности обучения: Мы сравнили время обучения моделей с и без преждевременной остановки. Модель с преждевременной остановкой достигла сопоставимой точности, потребовав значительно меньше оценщиков, что привело к более быстрому обучению.

Общее время выполнения скрипта: (0 минут 3.542 секунды)

Launch binder
Launch JupyterLite

Download Jupyter notebook: plot_gradient_boosting_early_stopping.ipynb

Download Python source code: plot_gradient_boosting_early_stopping.py

Download zipped: plot_gradient_boosting_early_stopping.zip

Связанные примеры

Преждевременная остановка стохастического градиентного спуска

Регрессия с помощью градиентного бустинга

Масштабирование параметра регуляризации для SVC

Сравнение случайных лесов и моделей градиентного бустинга с гистограммами

© 2007–2025 The scikit-learn developers
Licensed under the 3-clause BSD License.
https://scikit-learn.org/1.6/auto_examples/ensemble/plot_gradient_boosting_early_stopping.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API