XGBoost (二)公式推导

XGBoost 训练的核心是逐棵构建决策树,而每棵树的训练需要解决两个关键问题:如何进行节点分裂,以及如何确定叶子节点的输出值。

在训练过程中,XGBoost 需要同时兼顾训练误差和模型复杂度。传统的 square error 或 friedman mse 等分裂增益计算方法,无法直接满足这种联合优化的需求。因此,必须基于损失函数的优化目标,推导出适用于 XGBoost 的分裂增益计算逻辑。

同样地,树构建完成后,叶子节点的输出值也不能通过简单规则确定。其最优解需结合具体损失函数的特性,通过公式推导得到相应的计算方法。

简言之,公式推导的目的就是明确:分裂增益的计算方法,以及叶子节点输出的最优计算方法。

1. 目标函数

一部分衡量模型在训练集上的拟合情况(训练误差),另一部分惩罚模型复杂度(正则项)。

  • \( \gamma T \):表示对叶子节点数量的惩罚,\(\gamma\) 越大说明分裂一次带来的复杂度更大,越希望模型更加简单。
  • \( \frac{1}{2}\lambda \|w\|^2 \) 表示对模型输出的惩罚。\(\lambda\) 越大说明越希望模型输出保守一些。

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

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

我们对损失函数在 \( t-1 \) 处进行泰勒二阶展开,通过损失函数在该处的损失、以及一阶导和二阶导信息,得到近似的损失表示:

  • \( g_{i} \) 表示样本的一阶导
  • \( h_{i} \) 表示样本的二阶导

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


我们希望知道新树 \(f_{t}\)如何如何输出,才能使得损失最小。为了能够得到新树的最优输出,我们需要将上面的目标函数转换到叶子节点角度表示的形式:

注意:上述公式中,\( w \) 表示叶子节点的输出值。


对 \( w \) 求导,并令导数等于 0,得到叶子节点最优输出值表示:

注意:\(\lambda\) 最初是损失函数中正则化项的系数,经过推导,最终体现在上公示中的分母部分。回顾前面例子,当 \(G_{j} = -10\)、\(H_{j}=0.01\)、\(\lambda = 1\),则 \(w_{j} = 9.9\),输出值显著降低,避免了输出爆炸式的极端情况。

此时,将叶子节点输出值公式代回到原公式中,得到在最优输出的情况下的损失表示:

再简化下上面公式的表示:

\( G_{i} \) 表示叶子结点上一阶导之和,\( H_{i} \) 表示叶子结点上的二阶导之和。

公式中:第一项表示这个叶子节点在最优输出值下能带来的损失下降。第二项表示这个叶子节点自身的复杂度成本。

  • 如果公式的值为 负值,说明该叶子节点带来的损失下降大于它增加的复杂度惩罚,对目标函数有利,有助于下降。
  • 如果公式的值为 正值,说明该叶子节点带来的损失下降不足以抵消复杂度惩罚,它的存在反而会让目标函数上升。

换句话说,这个值衡量了每个叶子节点对目标函数的贡献。如果一个叶子节点能带来的损失下降(第一项)不足以抵消它的复杂度成本(第二项),那么这个叶子的存在会让整体目标函数变大,不应该被保留。

2. 分裂增益

单个叶子节点对目标函数的贡献(越小越好)

当某个节点分裂后,原节点的对目标函数的贡献就被替换为:新左节点的贡献 + 新右节点的贡献。

我们可以利用分裂前后的的差值来作为 XGBoost 的分裂增益计算方法:

  • 当 \( Gain \gt 0\),分裂前能降 -500,分裂后能降 -800,分裂后损失能多降 300,该分裂被考虑
  • 当 \( Gain \leq 0\),分裂前能降 -500,分裂后能降 -300,分裂后损失能少降 200,该分裂不考虑
  • Gain 越大说明,此次分裂的价值越大

这个公式就是 XGBoost 算法的分裂增益计算公式。如果一次分裂带来的损失下降不足以抵消它的复杂度成本,此次分裂不会被考虑。