梯度提升何时该停:用验证误差比较早停与完整训练

梯度提升把多个弱学习器组合成更强的预测模型,常见的弱学习器是决策树。模型一轮一轮地构建,每一轮新增的树都试图修正之前模型留下的误差。树越多,训练集上的误差往往越低,但这并不意味着对未见数据的预测会一直改善。

早停为训练设置了一条基于验证表现的停止规则:先留出一部分数据,用 validation_fraction 指定其比例;随后观察模型增加树时,在这份内部验证集上的损失如何变化。当连续若干轮没有取得超过容差 tol 的足够改善,就结束训练,连续观察的轮数由 n_iter_no_change 指定。训练完成后,实际使用的树数可以从 n_estimators_ 读取。

这是一种在泛化表现与计算开销之间取得平衡的办法。它不会保证找出全局最优树数,也不保证在任何数据上都降低误差。下面保留官方示例的完整流程:相同参数上限下,分别训练启用和未启用早停的模型,再对比误差曲线、树数和训练时间。

原文作者:The scikit-learn developers。依据 scikit-learn 1.9.1 官方实例翻译,核验日期为2026年10月5日。代码与原始图遵循 BSD-3-Clause,本文中的“编辑补充”用于说明实验边界。

准备数据

示例加载 California Housing Prices 数据集,取前600行,再按80%与20%的比例划分成训练数据和外部验证数据。random_state=42 固定这次划分的随机状态。

# 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)

编辑补充:fetch_california_housing() 可能需要联网下载数据,并在本地缓存。data[:600] 只是按现有顺序截取前600行,不是随机抽样600行,因此不能把这个小子集的表现当成整个加州房价数据集的代表性结论。

训练两个模型,记录成本

两个 GradientBoostingRegressor 共享最大树数1000、树深度5、学习率0.1和随机状态42。gbm_full 按设定完成全部迭代;gbm_early_stopping 则从传给 fit 的训练数据中留出10%,在连续10轮没有足够改善时停止。代码分别记录训练耗时与实际树数。

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_
600行数据分成480行外层训练和120行外部评估;早停又从480行中保留48行作内部验证,剩余432行拟合。
编辑原创示意图。外部 X_val 用于比较,内部保留集才决定什么时候停止。

编辑补充:这一划分意味着完整模型用480行拟合,而早停模型实际用432行拟合,另有48行用于内部停止判断。原文中的 X_val 有120行,它没有传入 fit。不要把“原文名为验证集的数据”与“估计器内部用来早停的数据”混为一谈。

计算每一轮的误差

staged_predict 会依次给出各个提升阶段的预测。示例分别对训练集和外部验证集计算均方误差(MSE),将结果写进四个列表,从而观察两个模型怎样收敛,而不是只比较最后一个数值。

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))

这里的训练曲线统一在整个 X_train 上计算。对早停模型而言,这480行还包含那48行内部验证样本,因此它并不是只针对“实际参与拟合的数据”计算的训练损失。两条训练曲线的差异,部分来自两种模型使用了不同的拟合数据。

把误差、时间和树数放在一起看

绘图包含三部分:左图显示训练误差随迭代变化,中图显示外部验证误差,右图用柱高表示训练时间,并在柱上标出实际树数。两张误差图采用对数纵轴,便于观察不同数量级的变化。

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()
原文三联图:完整训练与早停的训练MSE、外部验证MSE,以及训练时间;树数分别为1000和119。
原文发布的结果图,来源:The scikit-learn developers,BSD-3-Clause。蓝线为完整训练,橙线为早停;这是上游示例输出,不是本次运行结果。

原图中,完整模型训练了1000棵树,早停模型在119棵树时停止。完整模型的训练误差仍继续下降,但外部验证误差已经趋于平稳;早停模型在使用更少树的情况下取得了相近的外部验证表现,并显著减少了该次运行的训练时间。原页面还记录了整个脚本运行时间2.720秒,这与柱图中的单次模型拟合时间不是同一个指标。

原文据此强调早停的两项价值:在验证表现不再改善时抑制继续拟合带来的过拟合风险,并减少无效迭代、提高训练效率。这里能支持的结论限于这次实验;具体停止轮数和加速比例会随版本、样本、硬件和系统负载变化,不能当作普遍保证。

把示例用于自己的实验

先保持内部停止集与外部评估集职责清晰,再选择 validation_fraction、n_iter_no_change 和 tol。若看过外部曲线后反复修改参数,这份外部数据也参与了模型选择,不能继续承担最终无偏评估的职责;此时应另留最终测试集。时间比较也应在相同环境下重复进行,而不是只根据一次 time.time() 差值判断。

本稿只静态检查了代码和图文对应关系,未下载数据、训练模型或复现性能数字。上面的代码保留原文逻辑,包括仅用于枚举、后续没有使用的变量 i;没有把示例暗改为另一套实验。

来源与许可

原文:Early stopping in Gradient Boosting;完整 Python 源码。版权为 © 2007–2026 The scikit-learn developers,BSD-3-Clause。完整版权声明、再分发条件及免责声明保留于下方 LICENSE.txt;不得用作者或贡献者名称为衍生产品背书。中文翻译和新增示意图与原始实验图已分别标明。

版权与许可全文

以下保留本页涉及的来源材料或示例代码的版权、许可条件与免责声明;各自适用范围依原声明。中文翻译及编辑标注:未完纪,2026-10-05。

LICENSE.txt

BSD 3-Clause License

Copyright (c) 2007-2026 The scikit-learn developers.
All rights reserved.

Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are met:

* Redistributions of source code must retain the above copyright notice, this
  list of conditions and the following disclaimer.

* Redistributions in binary form must reproduce the above copyright notice,
  this list of conditions and the following disclaimer in the documentation
  and/or other materials provided with the distribution.

* Neither the name of the copyright holder nor the names of its
  contributors may be used to endorse or promote products derived from
  this software without specific prior written permission.

THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
© 版权声明
THE END
喜欢就支持一下吧
点赞0 分享
评论 抢沙发

请登录后发表评论

    暂无评论内容