GBDT 梯度提升树(二)公式推导

在 GBDT 的训练过程中,每棵新树的核心任务是拟合前一轮模型的负梯度,以此逐步修正预测误差。要完成这棵树的训练,需明确两个关键问题:

  1. 树的分裂规则:如何计算特征分裂时的增益,从而选择最优分裂点?
  2. 叶子节点输出:分裂结束后,每个叶子节点应输出什么值,才能最大化误差修正效果?

先看第一个问题:GBDT 训练每棵树的核心目标是精准拟合负梯度、降低整体误差,不需要考虑树的复杂度。因此,选择能直接反映拟合效果提升的分裂增益指标即可,常用的有平方误差(squared_error)或弗里德曼均方误差(friedman_mse),scikit-learn 等工具默认采用 friedman_mse,这一步无需额外推导公式,直接选用成熟指标即可。

再看第二个问题:叶子节点输出值的确定,这是 GBDT 训练的关键难点,且结果与损失函数直接相关,不能一概而论:

  • 若使用平方损失函数,情况较简单,叶子节点的最优输出值,等价于该节点内所有样本负梯度的均值。
  • 若使用非平方损失函数(如 log loss),则无法直接用负梯度均值作为最优输出。此时必须通过最小化该叶子节点内所有样本的累计损失来求解最优值 。
  • 这一步是需要推导的。

简单来说:分裂增益的选择是 “选工具”,优先挑能高效拟合负梯度的成熟方法。而叶子节点输出是 “定结果”,必须贴合损失函数的特性,平方损失有简洁结论,非平方损失则需通过损失最小化单独求解。

1. 目标函数

公式表示,所有训练样本的损失总和。

公式表示,模型的预测结果 \( \hat{y} \) 等于 \( m-1 \) 树的输出结果 + 当前树 \(f_{m}\)的输出结果。

我们通过一个例子来理解损失函数:假设目标值为 100、初始预测值为 20,此时初始损失为 6400。后续每新增一棵树,都会通过拟合负梯度的方式,最终实现损失函数的最小化。

轮次拟合残差实际输出总预测更新当前损失损失下降幅度说明
第 1 棵树8030502500-3900损失大幅下降
第 2 棵树5010601600-900损失持续下降
第 3 棵树403090100-1500损失进一步下降

后续树循环拟合剩余残差,总预测逐步逼近目标值,总损失持续向 0 收敛。


优化目标:确定每棵新树 \(f_{m}\) 如何输出(叶子节点输出值),能够最小化损失函数。

2. 泰勒展开

由于 \( l \) 可能是任意复杂的非线性函数,并且我们要优化的 \( f_{m}(x) \) 是整个损失函数内部的一部分,直接对 \( f_{m}(x) \) 优化不现实。所以,我们使用泰勒二阶展开将损失函数近似成一个关于 \( f_{m}(x) \) 的二次函数,这样就能继续优化 \(f_{m}\)。

泰勒二阶展开是利用损失函数在某个已知点、以及其一阶导和二阶导信息,近似表示加入新树后损失表示。泰勒二阶展开公式如下:

对应到当前的场景:

  • \(a\) 表示截止到 m-1 树时的模型输出值
  • \(x\) 表示截止到 m 树时的模型输出值
  • \(x-a\) 表示第 m 棵树带来的模型输出增量
  • \(f(a)\) 表示截止到第 m-1 棵树时,模型的损失函数值
  • \(f^{\prime}(a)\) 表示损失函数在 a 处对模型输出的一阶导数
  • \(f^{\prime\prime}(a)\) 损失函数在 a 处对模型输出的二阶导数

我们对损失函数在已知点 \( m-1 \) 位置进行泰勒二阶展开:

  • \( g_{i} \) 表示第 i 样本在损失函数上的一阶导值
  • \( h_{i} \) 表示第 i 样本在损失函数上的二阶导值
  • \( f_{m}(x_{i}) \) 表示样本 \( x_{i} \)在第 m 棵树的输出值

\( l(y_i, F_{m-1}(x_i)) \) 表示还没有加入新树时的损失,是常数项,可以去掉,我们就得到了一个关于 \( f_{m}(x_{i}) \) 的公式:

我们通过最小化这个目标函数,来找到第 m 棵树的最优输出值(叶子节点的输出值)。


我们要推导叶子结点最优输出值,将上面公式转为从叶子节点角度表示的目标函数。为了简便,我们用 \( \gamma_{j} \) 表示第 \(j\) 个叶子节点的输出值。

每个样本的输出值对应某个叶子节点的值,所以我们可以以叶子节点为单位计算每个样本的输出值、一阶导、二阶导,那么得到如下公式:

我们将公式再简化一些:

3. 最优输出

叶子结点如何输出,目标函数最小?我们只需要对 \( \gamma \) 求导,并令导数等于 0。


假设我们使用 log loss、square loss 损失函数,那么叶子节点的输出值公式为:

square loss

注意:\( (y_i – \hat{y}_i^{(m-1)}) \) 表示本轮样本的负梯度。

log loss