中文

用于斜决策树的普通梯度下降法

机器学习 2024-10-21 v3 人工智能

摘要

决策树(DTs)是主要的强非线性人工智能模型之一,因其在表格数据上的高效性而受到重视。然而,学习准确的决策树,尤其是斜决策树,是复杂的,并且确实需要大量的训练时间。此外,决策树容易过拟合,例如,它们在回归任务中众所周知地“泛化能力差”。最近,一些工作提出了使(斜)决策树可微的方法。这使得能够使用高效的梯度下降算法来学习决策树。它还通过在叶节点学习回归器与树中决策同步进行,从而赋予模型泛化能力。先前使决策树可微的方法要么依赖于树内部节点的概率近似(软决策树),要么依赖于内部节点梯度计算的近似(量化梯度下降)。在这项工作中,我们提出了 DTSemNet,这是一种新颖的、语义等价且可逆的编码方式,将(硬、斜)决策树编码为神经网络,并使用标准的普通梯度下降法。在各种分类和回归基准上的实验表明,使用 DTSemNet 学习的斜决策树比使用最先进技术学习的相似大小的斜决策树更准确。此外,决策树的训练时间显著减少。我们还通过实验证明,在具有物理输入(维度 32\leq32)的强化学习设置中,DTSemNet 可以像神经网络策略一样高效地学习决策树策略。代码可在 https://github.com/CPS-research-group/dtsemnet 获取。

关键词

引用

@article{arxiv.2408.09135,
  title  = {Vanilla Gradient Descent for Oblique Decision Trees},
  author = {Subrat Prasad Panda and Blaise Genest and Arvind Easwaran and Ponnuthurai Nagaratnam Suganthan},
  journal= {arXiv preprint arXiv:2408.09135},
  year   = {2024}
}

备注

Published in European Conference on Artificial Intelligence (ECAI), 2024. Full version (includes supplementary material)