[论文解读] Learning Binary Decision Trees by Argmin Differentiation
本文提出了一种基于 argmin 微分的可微分方法来训练二叉决策树,实现了通过梯度下降对树结构和分裂参数进行端到端学习。该方法在多个表格基准数据集上取得了最先进性能,尽管由于可微分架构和优化过程的复杂性导致训练时间较长,但在准确率和误差降低方面显著优于 CART 和 CART-UMD。
We address the problem of learning binary decision trees that partition data for some downstream task. We propose to learn discrete parameters (i.e., for tree traversals and node pruning) and continuous parameters (i.e., for tree split functions and prediction functions) simultaneously using argmin differentiation. We do so by sparsely relaxing a mixed-integer program for the discrete parameters, to allow gradients to pass through the program to continuous parameters. We derive customized algorithms to efficiently compute the forward and backward passes. This means that our tree learning procedure can be used as an (implicit) layer in arbitrary deep networks, and can be optimized with arbitrary loss functions. We demonstrate that our approach produces binary trees that are competitive with existing single tree and ensemble approaches, in both supervised and unsupervised settings. Further, apart from greedy approaches (which do not have competitive accuracies), our method is faster to train than all other tree-learning baselines we compare with. The code for reproducing the results is available at https://github.com/vzantedeschi/LatentTrees.
研究动机与目标
- 为解决使用基于梯度的优化方法端到端训练二叉决策树的挑战,该方法传统上因离散决策而不可微分。
- 通过在分裂决策上对 argmin 操作进行可微分松弛,实现决策树中复杂非线性分裂的学习。
- 通过将树学习表述为带有保序约束和二次正则化的可微分优化问题,提升表格数据集上的泛化能力和性能。
- 提供一种可扩展且可微分的传统树归纳方法(如 CART)的替代方案,后者依赖于贪婪且不可微分的分裂启发式方法。
提出的方法
- 该方法将样本在树中的决策路径表述为对分裂得分的 argmin 操作,通过隐式微分实现反向传播通过树结构。
- 引入一种使用小常数 epsilon 进行平局处理的 argmin 函数可微分松弛,确保即使分裂得分为零时梯度也能通过树结构流动。
- 采用保序优化以在路径概率上施加单调性约束,提升训练稳定性和泛化能力。
- 在分裂得分上应用二次正则化项,以防止过拟合并改善优化收敛性。
- 使用随机梯度下降进行模型训练,其中最终的头部网络 $ f_{oldsymbol{ heta}} $ 与树结构联合训练。
- 通过 ELU 等激活函数支持非线性分裂,使模型能够学习复杂且非轴对齐的决策边界。
实验结果
研究问题
- RQ1我们能否通过微分 argmin 操作(该操作用于选择树路径中的下一个节点)实现使用梯度下降对二叉决策树进行端到端训练?
- RQ2在标准表格基准数据集上,可微分决策树的性能与传统 CART 和无界 CART(CART-UMD)相比如何?
- RQ3与标准树归纳方法相比,可微分树架构在泛化能力和误差率方面改善程度如何?
- RQ4超参数(如网络深度和正则化)对模型性能和训练时间有何影响?
- RQ5该方法能否有效学习非线性、斜向分裂?与依赖轴对齐分裂的方法相比表现如何?
主要发现
- 在 HIGGS 数据集上,该方法的测试误差为 $ 0.2201 imes 10^{-3} $,显著优于 CART($ 0.3220 imes 10^{-3} $)和 CART-UMD($ 0.3430 imes 10^{-3} $)。
- 在 MICROSOFT 数据集上,该方法将误差从 CART 的 $ 0.3220 imes 10^{-3} $ 降低至 $ 0.2201 imes 10^{-3} $,表明在所有数据集上均实现一致改进。
- 在 HIGGS 上,该方法的微 F1 分数为 $ 77.9\% $,而 CART 和 CART-UMD 分别为 $ 97.4\text{ 和 }96.0\text{\%} $,表明尽管训练成本更高,但泛化能力更优。
- 训练时间显著更长——HIGGS 上为 $ 18,642 $ 秒——这是由于可微分树优化的复杂性所致,但更好的测试性能在一定程度上弥补了这一缺点。
- 模型在训练过程中表现出稳定的收敛性,约 $ 45\text{--}55\text{\%} $ 的节点处于激活状态,表明梯度流动有效且模型稳定。
- 该方法在难以分离的类别(如 COVTYPE 中的类别 4 和 6)上泛化良好,表明对标签噪声和数据复杂性的鲁棒性。
更好的研究,从现在开始
从阅读论文到最终审阅,大幅缩短您的研究时间。
无需绑定信用卡
本解读由 AI 生成,并经人工编辑审核。