Paper Reading: Differentiable Decision Tree via "ReLU+Argmin" Reformulation
Paper Reading 是从个人角度进行的一些总结分享,受到个人关注点的侧重和实力所限,可能有理解不到位的地方。具体的细节还需要以原文的内容为准,博客中的图表若未另外说明则均来自原文。
| 论文概况 | 详细 |
|---|---|
| 标题 | 《Differentiable Decision Tree via “ReLU+Argmin” Reformulation》 |
| 作者 | Qiangqiang Mao, Jiayang Ren, Yixiu Wang, Chenxuanyin Zou, Jingjing Zheng, Yankai Cao |
| 发表会议 | Thirty-Ninth Conference on Neural Information Processing Systems(NeurIPS 2025) |
| 发表年份 | 2025 |
| 会议等级 | CCF-A |
| 论文代码 | https://github.com/YankaiGroup/RADDT |
作者单位:
- University of British Columbia, Vancouver, Canada
研究动机
决策树因其可解释性强、结构轻量而受到机器学习领域的广泛关注,尤其适用于需要明确决策过程的任务。但是传统决策树在分支节点进行硬性分割,在决策路径上做唯一性选择,其硬决策的特征导致整个模型是不可微的。这限制了决策树在需要梯度反向传播的优化框架中的应用,例如将其嵌入到强化学习策略、神经网络或其他端到端的可微系统中。在实际应用中,决策树(特别是经典的正交树)的测试精度常常低于其他复杂模型(如神经网络、集成方法),这迫使使用者在可解释性和高精度之间做出取舍。
使用线性组合特征的斜决策树在数据边界为超平面时,可能以更小的树深获得更高的精度。但其训练更具挑战性,因为每个节点的分割有无穷种线性组合可能。现有方法及其缺陷总结如下:
| 现有方法 | 局限性 |
|---|---|
| 贪心算法 | 如 CART, OC1 使用逐节点优化,可能导致后续节点的分割质量下降,陷入局部最优。 |
| 全局优化方法 | 如基于混合整数规划的 MIP 方法,虽然能优化整棵树,但计算成本极高,可扩展性差,难以处理大规模数据或回归任务。 |
| 直通估计器 | 可能忽略关键梯度信息,导致学习效果不佳。如 GradTree 在原文中报告其在 36 个数据集的 17 个上表现不如 CART。 |
| 软决策树 | 使用 Sigmoid 等函数将硬决策软化,引入概率性路径和预测。这虽然可微,但改变了决策树的本质,在某些需要硬决策和确定性预测的场景中并不适用,如显式模型预测控制中的分段仿射控制律。 |
基于以上背景,该论文旨在同时解决决策树的两个瓶颈:
- 如何使硬决策树变得可微?设计一种不引入软决策的、保留经典硬决策树本质(硬分叉、唯一决策路径、确定性预测)的可微重构方法,使其能够集成到基于梯度的优化框架中。
- 如何提升决策树的测试精度?通过一种可扩展的、全局优化整棵树的方法,找到分割质量更高、预测更准确的斜决策树,使其精度不仅能超越传统贪心算法,甚至能够与强大的树集成方法竞争。
文章贡献
针对经典决策树的不可微性和较低测试精度两个瓶颈,本文提出了一种名为 RADDT 的可微硬斜决策树模型。其核心是通过 ReLU + Argmin 的数学重构,将不可微的硬决策树训练问题转化为一个可微的无约束优化问题。首先,本文提出了 ReLU+Argmin 的重构方法,用 ReLU 函数量化样本在决策路径上的违规值,并用 Argmin 操作识别唯一零违规路径,从而在不引入任何软决策的前提下,为硬决策树建立了一个精确的数学表达。接着,为了解决重构中 Argmin 不可微的问题,采用缩放 Softmin 函数在训练时近似,并设计了多轮热启动退火策略来平衡近似精度与优化稳定性。最终,基于此重构构建了名为 RADDT 的可微决策树训练框架,实现了对整个决策树参数的梯度优化。大量实验表明,该方法不仅测试精度显著超越了包括 CART、ORT-LS 在内的多种基线决策树,甚至能与 RF、XGBoost 等集成方法媲美,并且展现出强大的可扩展性。
斜决策树训练的无约束优化重构
本文旨在将经典的、不可微的硬决策树训练问题,精确地重构为一个可微的、无约束的优化任务,从而能够利用梯度下降等优化工具来求解。
基于 ReLU+Argmin 的硬决策树重构
该步骤用于消除软决策树在分支和预测时的“软性”,重构一个保持硬分叉和确定性决策路径的可微决策树。首先定义基于 ReLU 的硬分叉与正确路径表述,主要使用 ReLU 函数来量化样本在每个分支节点是否遵循正确方向的违规程度。对于一个要到达特定叶子节点 \(t\) 的样本 \(x_i\),其必须经过一系列祖先节点 \(A_t\),其中 \(A_t^l\) 是左转祖先集合,\(A_t^r\) 是右转祖先集合。在左转祖先 \(j \in A_t^l\),样本必须满足 \(a_j^T x_i \le b_j\) 才能向左。其违规值 \(v_{i,j}\) 定义为:
若 \(v_{i,j} = 0\),表示方向正确;否则 \(v_{i,j} > 0\),表示方向错误。在右转祖先 \(j \in A_t^r\),样本必须满足 \(a_j^T x_i > b_j\) 才能向右。其违规值定义为:
与软决策树在 \((0,1)\) 之间输出分支概率不同,这里的 ReLU 输出是确定性的零或非零值,保留了硬分支的本质。接着通过 Argmin,给出累计违规与唯一决策路径表述。累计违规定义为样本 \(x_i\) 到达叶子节点 \(t\) 的总违规值,是其在所有祖先节点上违规值之和:
\(U_{i,t} = 0\) 意味着样本 \(x_i\) 完全遵循了到达叶子 \(t\) 的所有正确方向。然后进行唯一路径识别,对于给定的 \(x_i\),有且仅有一个叶子节点 \(t\) 满足 \(U_{i,t} = 0\),这条零违规路径就是其唯一的决策路径。用 Argmin 操作来识别:
其中 \(U_i = \{U_{i, \lfloor T/2\rfloor+1}, \dots, U_{i, T}\}\),\(M(\cdot)\) 是一个 one-hot 编码函数。使得样本分配和预测是确定性的,而非概率性的。
论文以深度为 2 的树为例,要到达叶子 5,样本必须在节点1左转 \((v_{i,1}=0)\),在节点 2 右转 \((v_{i,2}=0)\),则 \(U_{i,5} = 0\)。对于其他叶子节点,累计违规值 \(U_{i,t} > 0\)。因此,\(M(U_{i,5}) = 1\) 表明 \(x_i\) 被唯一地分配到叶子 5。

基于 ReLU+Argmin树训练的无约束优化
基于上述重构,可以将整个决策树的训练目标 \(\mathcal{L}\) 表示为:
其中,\(\ell(\cdot)\) 是损失函数,回归用均方误差,分类用交叉熵。ReLU 在零点处的梯度行为已有大量研究,通常认为其对梯度优化影响可忽略不计,这已被 ReLU 网络的成功所证明。Argmin 则是一个离散的、不可微的操作,是梯度回传的主要障碍。为了解决这个问题,论文在训练阶段使用一个缩放 Softmin 函数 \(S(\cdot)\) 来近似 Argmin 的行为:
其中 \(\alpha\) 为缩放因子,\(\alpha\) 越大则 Softmin 的输出越接近理想的 one-hot Argmin 输出,即零违规路径对应的值接近 1 其他接近 0,近似越精确。然而,过大的 \(\alpha\) 会导致梯度“陡峭”,引发数值不稳定,从而影响优化。选择 \(\alpha\) 以平衡近似精度和优化稳定性是一个关键挑战。
将近似后的 Softmin 代入,得到最终可微的无约束优化目标 \(\mathcal{L}\):
需要注意的是近似仅在训练时使用,即 Softmin 仅用于训练时回传梯度。在推理阶段,模型会使用精确的 Argmin 来确定路径,并使用确定性方法(根据分配到叶子的样本计算平均值或进行线性回归)重新计算叶子的预测值 \(\theta\),以保证决策的确定性。因为零违规值 \(U_{i,t}=0\) 和对应的 ReLU 硬决策是在 Softmin 近似之前就已经被确定下来的。Softmin 只是为了让这个确定的、硬决策的结构能够被梯度优化,而不是引入软决策。
基于梯度的整树优化之可微决策树训练
此处阐述了如何高效、稳定地求解无约束优化问题,核心在于解决缩放 Softmin 中近似精度与数值稳定性的矛盾,并设计一个完整的梯度优化框架。
用于缩放 Softmin 的多轮热启动退火策略
缩放因子 \(\alpha\) 控制 Softmin 对 Argmin 的近似程度。\(\alpha\) 越大近似越精确,但梯度越“陡峭”导致优化越不稳定;\(\alpha\) 越小则梯度越平滑,但近似越不精确。难以选取一个固定的 \(\alpha\) 值来同时保证最优性和稳定性。解决思路是采用退火思想,但将其应用于多轮独立的优化任务之间,而非单个训练过程中调整。
具体操作为:在一个预设的范围 \([\alpha_{\min}, \alpha_{\max}]\) 内,以对数空间采样一组递增的缩放因子 \(\{\alpha_1, ..., \alpha_m\}\)。然后用最小的缩放因子 \(\alpha_1\) 启动优化任务,此时问题相对平滑,容易收敛得到一个初步解。然后,将上一个任务(使用 \(\alpha_k\))得到的优化解,即树的分割参数 \(A, b\) 和叶子参数 \(\theta\) 作为热启动的初始点,用于下一个使用更大缩放因子 \(\alpha_{k+1}\) 的优化任务。
这种方法允许优化过程逐渐适应越来越精确(但也越来越尖锐)的近似形式,用平滑近似下的解去初始化尖锐近似的优化,能有效缓解直接使用大 \(\alpha\) 可能导致的数值不稳定和优化困难,简单的二分搜索策略是孤立地测试不同的 \(\alpha\) 值,无法利用前一个 \(\alpha\) 的优化结果来热启动下一个。实验证明,本文的多轮热启动退火策略在训练最优性上平均优于二分搜索策略。

初始值及其调整对梯度优化的影响
初始解的质量对梯度优化至关重要,在 Softmin 近似下如果样本非常接近决策边界(即 \(a_j^T x_i - b_j \approx 0\)),其违规值会是一个接近零的正数。虽然这不影响精确的 Argmin 判断(它只认严格的 0),但会导致 Softmin 无法理想地输出接近 1 的值,从而降低近似精度和影响优化。解决方案是调整初始的分割参数 \((A, b)\),以最大化决策边界与最近样本之间的间隔,同时不改变当前分割下的样本分配。
可以将每个分支节点的分割视为一个二分类问题(左/右)。通过计算边界两侧最近样本的 \(a_j^T x_i\) 的中值,或训练一个线性 SVM 来获得一个最大间隔超平面,然后使用该超平面的权重和偏置来更新 \(a_j\) 和 \(b_j\)。极端情况是当初始化出现 \(a_j = 0\) 且 \(b_j = 0\) 时,此调整策略会失效。不过,作者指出实践中这种情况极为罕见,且通过多重随机初始化可以很大程度上规避其负面影响。
可微决策树优化框架
将上述策略整合为一个系统的优化框架 RADDT,其核心特点在于整树同时优化:与贪心算法逐节点优化不同,RADDT 同时优化所有分支节点的分割参数 \((A, b)\) 和所有叶子节点的预测参数 \((\theta)\)。算法流程如下:
- 多重随机初始化:执行 \(N_{\text{start}}\) 次独立的优化过程,每次从不同的随机初始解开始,以增加找到全局更优解的概率。
- 单次优化过程:
a. 初始解调整:对当前随机初始化的 \((A, b)\) 应用间隔最大化方法调整。
b. 多轮退火优化:在一系列递增的缩放因子 \(\{\alpha_1, ..., \alpha_m\}\) 上顺序执行梯度下降优化,并使用热启动。
c. 候选树评估:在每一轮退火(即每个 \(\alpha\))优化结束后,确定性地(不使用 Softmin 近似)重新计算叶子预测参数 \(\theta\),并计算精确的损失 \(\mathcal{L}\)。 - 最优树选择:从所有随机初始化和所有退火轮次产生的候选树中,选择精确损失 \(\mathcal{L}\) 最小的那棵树作为最终模型。
整个框架可以使用现有深度学习工具(如 PyTorch)实现,并能利用其自动微分和 GPU 加速功能。
超参数分析
作者分析了框架中的几个关键超参数,并指出其影响直观,通常无需复杂调参:
| 超参数 | 分析 |
|---|---|
| 多启动数 \(N_{\text{start}}\) | 增加此值可提升找到更优解的概率,但会增加计算成本。实验表明,增加 \(N_{\text{start}}\) 对所有深度树的训练最优性都有提升,尤其在浅层时更明显。 |
| 训练轮数 \(N_{\text{epoch}}\) | 增加轮数可提高训练精度,但也增加耗时。 |
| 缩放因子范围 \([\alpha_{\min}, \alpha_{\max}]\) | 设置为 \([2, 200]\) 即可满足需求,确保小值梯度平滑,大值近似精确。 |
| 采样因子数量 | 采样更多 \(\alpha\) 值(如 5 个)可使退火过程更稳定,但会增加计算迭代次数。 |
| 学习率 \(\eta\) | 使用 PyTorch 的标准学习率调度器(如带热重启的余弦退火)即可,初始值设为 0.01。 |

实验结果
数据集和实验设置
论文主要评估两种树:
| 模型 | 说明 |
|---|---|
| RADDT | 带常数(均值)预测的斜决策树 |
| RADDT-Linear | 带线性预测的斜决策树 |
基线方法包含 14 种,涵盖不同类型,以确保对比的广泛性:
| 算法类型 | 对比模型 |
|---|---|
| 贪心算法 | CART, HHCART, RandCART, OC1 |
| 非贪心算法 | TAO |
| 梯度优化树 | GradTree, DGT, DTSemNet, SoftDT, LatentTree |
| 局部搜索 | ORT-LS |
| 集成方法 | TEL(梯度树集成), RF, XGBoost |
实验在四组数据集上进行:
| 数据集组 | 特点 | 说明 | 目的 |
|---|---|---|---|
| Group(i) | 中型回归 | 17 个数据集(<41k 样本) | 实验的主要焦点 |
| Group(ii) | 基线原用数据集 | 包括 GradTree 的 27 个分类数据集、软决策树的 4 个回归数据集,以及 DGT-Linear/TAO-Linear/DTSemNet 共享的 5 个回归数据集 | 用于更可信的对比 |
| Group(iii) | 小/合成数据集 | 4 个约 100 个样本的真实数据集(用于对比全局最优 MIP 解)和 3 个 5000 样本的合成数据集(已知真实情况) | 用于分析训练最优性 |
| Group(iv) | 大规模数据集 | 7 个百万级样本数据集 | 用于评估可扩展性 |
回归任务主要使用决定系数 \(R^2\),分类任务使用宏 F1 分数进行评估。同时使用 Friedman 排名对方法进行统计排序,排名越低越好。
对比实验
在 Group (i) 数据集上与其他决策树的对比,RADDT 在所有决策树基线中取得了最优的平均测试精度。其平均 \(R^2\) 为 82.39%,显著优于基础方法 CART(74.85%)、贪心斜树 HHCART(76.75%)和局部搜索 ORT-LS(78.67%)。Friedman 排名为 1.76,位居第一。GradTree、SoftDT 表现不佳,DGT 和 LatentTree 由于在深度 12 时可扩展性不足,在本组固定深度对比中未列出。

在 Group (i) 数据集上与树集成方法的对比,可见 RADDT-Linear 良好,甚至超越了强大的树集成方法。RADDT-Linear 平均 \(R^2\) 为 84.63%,表现优于随机森林 RF(82.62%)、XGBoost(83.51%)和梯度树集成 TEL(83.87%)。Friedman 排名与 XGBoost 并列最高(2.24)。比较单一树与集成方法本身并不公平,任何基学习器都可通过集成提升。此对比意在表明,优化的单一树可以达到媲美集成的精度,且参数更少。

基于 Group (ii) 数据集进行复现了对比,在 27 个分类数据集上,RADDT 的宏 F1 分数平均比 GradTree 高出 4.99%。在 4 个回归数据集上,RADDT 的 \(R^2\) 平均比软决策树高出 6.88%。在 5 个共享回归数据集上,RADDT-Linear 在 4 个上表现最优,平均排名 1.6,优于DTSemNet(2.4)、TAO-Linear(2.6)和 DGT-Linear(3.2)。对表 1 结果进行了配对T检验。所有对比的 p 值均小于0.1,且 t 统计量为正,从统计学上证实了 RADDT 显著优于其他决策树方法。

训练最优性分析
通过固定深度(D=2,4,8,12)训练,分析模型在训练集上的表现,以探究其获得高测试精度的原因。RADDT 在训练精度(\(R^2\))上显著优于其他方法,表明其优化算法能够找到训练损失更低的解,这直接转化为了更高的测试精度。在深度 2、4、8 时,训练 \(R^2\) 分别比 ORT-LS 高出 5.35%、2.66%、0.02%。在所有深度,训练 \(R^2\) 平均比 CART 大幅领先,如深度 2 时高出 24.83%。平均训练精度显著高于 GradTree 和 SoftDT。在可行深度上,平均训练精度高于 LatentTree 10.39%,高于 DGT 7.52%。在 4 个极小数据集上,全局 MIP 方法(ORT-MIP)仅能求解深度 2 的树。RADDT 的解与全局最优解的差距仅为 2.82%,表明其能接近全局最优。在 3 个由决策树规则生成的合成数据集上,RADDT 的训练/测试精度与真实值(100%)的平均差距仅为 0.64%,进一步验证了其优化能力。

可扩展性评估
在基于 Group (iv) 百万级数据集上,RADDT 展现了可扩展性,能够在百万级数据集上成功训练深度为 12 的树,而多数梯度优化基线在此规模下失败。RADDT 在深度 12 上仍稳定工作平均训练/测试 \(R^2\) 分别比 CART 高 4.83% 和 3.15%。GradTree 和 SoftDT 无法产出有效解,DGT 和 LatentTree 在深度较大时(如8, 12)无法运行。在它们可运行的较浅深度(2, 4),RADDT 在精度和时间上均具优势,例如深度 4 时训练速度快 DGT 42 倍。

RADDT 的训练时间与样本数 \(n\)、特征数 \(p\) 和深度 \(D\) 成线性关系,复杂度为 \(O(n \cdot p \cdot (2^D - 1))\)。实现支持多 GPU 并行,矩阵运算主导的计算模式使其能充分利用硬件加速。在 ailerons数据集上,深度 12 时 RADDT 比可扩展性第二好的方法 ORT-LS 快 432 倍。

模型复杂度与推理时间
斜决策树在相同深度下比正交树(如 CART)参数更多,但它能以更浅的深度达到相同或更高的精度,从而在总体上可能使用更少的参数。结果可见要达到 CART 74.85% 的精度,RADDT 仅需平均深度 2.82、参数 101.71 个,比 CART 少 31 倍。RADDT-Linear 仅需深度 1、参数 40.47 个,少 78 倍。在追求最高精度时,RADDT(深度 7.29)和 RADDT-Linear(深度 4.82)的参数数会超过 CART,但这是为换取更高精度(82.39%, 84.63%)付出的代价。因参数更多,推理略慢于 CART。但在精度与 CART 持平时,因其深度和节点数更少,推理时间与 CART 相当甚至更快。

可解释性分析
在相同深度下,正交树(如 CART)通常比斜树更易解释。但当斜树能以浅得多的深度达到相同或更高精度时,其整体可解释性可能更强。与所有决策树一样,RADDT 能提供清晰的 IF-THEN 规则链,每个样本有确定路径便于追踪和解释。虽然斜分割的线性组合比单特征阈值复杂,但极浅的树深大大降低了整体规则的复杂度。以“混凝土抗压强度”数据集为例,展示了一个深度仅为 3 的 RADDT-Linear 树,其规则虽为线性组合,但整体非常简洁易懂。

消融实验
本文验证了两个核心策略的有效性。相比不使用任何缩放(\(\alpha=1\)的标准Softmin),多轮热启动退火策略平均提升训练精度 7.1%。在已有退火策略的基础上,初始化调整策略进一步带来平均 1.8% 的训练精度提升。

局限性分析与未来工作
论文指出了当前方法的四点主要局限,本文的未来工作包括:设计更好的正则化;将本文的优化策略(如初始化、退火)应用于其他梯度树(如 SoftDT);探索用 Entmax 函数替代 Softmin 等。
| 局限性 | 说明 |
|---|---|
| 正则化不足与过拟合 | 在深度较大时(如 12 层)出现明显过拟合。初步的 L1 正则化有改善,但设计更有效的正则化策略是未来方向。 |
| 理论分析欠缺 | 多轮热启动退火策略的有效性目前主要基于实证,缺乏严格理论分析。 |
| 初始化调整的角落情况 | 当 \(a_j=0\) 且 \(b_j=0\) 时策略失效,虽然这种情况罕见,但仍属理论缺陷。 |
| 与集成的对比深度 | 论文的集成方法对比已做充分调参,但更极致的正则化调参可能进一步提升集成方法性能。 |
优点和创新点
个人认为,本文有如下一些优点和创新点可供参考学习:
- 提出了 ReLU+Argmin 的硬决策树精确可微重构,在不牺牲决策树硬分叉和确定性路径核心特质的前提下,实现了决策树基于梯度的优化。
- 针对重构中的不可微难题,设计了多轮热启动退火策略,有效解决了近似精度与数值稳定性之间的矛盾,是优化成功的关键。
- 所提出的 RADDT 框架在测试精度上超越了现有主流决策树,并具备了在大规模数据集上训练深层斜决策树的强大可扩展性。

浙公网安备 33010602011771号