Differentiable Decision Tree via "ReLU+Argmin" Reformulation
Qiangqiang Mao, Jiayang Ren, Yixiu Wang, Chenxuanyin Zou, Jingjing Zheng, Yankai Cao
Abstract
Decision tree, despite its unmatched interpretability and lightweight structure, faces two key issues that limit its broader applicability: non-differentiability and low testing accuracy. This study addresses these issues by developing a differentiable oblique tree that optimizes the entire tree using gradient-based optimization. We propose an exact reformulation of hard-split trees based on “ReLU+Argmin” mechanism, and then cast the reformulated tree training as an unconstrained optimization task. The ReLU-based sample branching, expressed as exact-zero or non-zero values, preserve a unique decision path, in contrast to soft decision trees with probabilistic routing. The subsequent Argmin operation identifies the unique zero-violation path, enabling deterministic predictions. For effective gradient flow, we approximate Argmin behaviors by scaling softmin function. To ameliorate numerical instability, we propose a warm-start annealing scheme that solves multiple optimization tasks with increasingly accurate approximations. This reformulation alongside distributed GPU parallelism offers strong scalability, supporting 12-depth tree even on million-scale datasets where most baselines fail. Extensive experiments demonstrate that our optimized tree achieves a superior testing accuracy against 14 baselines, including an average improvement of 7.54% over CART .
Ask about this paper
Your agent reads all of it.
Lune indexed this paper to the last equation, along with the top-tier papers that cite it. Ask a question and the answer quotes them.
Builds on10
- Generalized and Scalable Optimal Sparse Decision TreesJimmy Lin, Chudi Zhong, Diane Hu, Cynthia Rudin et al.ICML 2020 · 174 citations
- The Tree Ensemble Layer: Differentiability meets Conditional ComputationHussein Hazimeh, Natalia Ponomareva, Petros Mol, Zhenyu Tan et al.ICML 2020 · 95 citations
- A Scalable MIP-based Method for Learning Optimal Multivariate Decision TreesHaoran Zhu, Pavankumar Murali, Dzung T. Phan, Lam M. Nguyen et al.NeurIPS 2020 · 47 citations
- Hierarchical Shrinkage: Improving the accuracy and interpretability of tree-based modelsAbhineet Agarwal, Yan Shuo Tan, Omer Ronen, Chandan Singh et al.ICML 2022 · 37 citations
- Smaller, more accurate regression forests using tree alternating optimizationArman Zharmagambetov, Miguel Á. Carreira-PerpiñánICML 2020 · 34 citations
Related papers
- Hinge Regression Tree: A Newton Method for Oblique Regression Tree SplittingHongyi Li, Han Lin, Jun XuICLR 2026 · 2 citations
- Oblique Decision Trees from Derivatives of ReLU NetworksGuang-He Lee, Tommi S. JaakkolaICLR 2020 · 25 citations
- GradTree: Learning Axis-Aligned Decision Trees with Gradient DescentSascha Marton, Stefan Lüdtke, Christian Bartelt, Heiner StuckenschmidtAAAI 2024 · 17 citations
- Learning Binary Decision Trees by Argmin DifferentiationValentina Zantedeschi, Matt J. Kusner, Vlad NiculaeICML 2021 · 16 citations
- Feature Learning for Interpretable, Performant Decision TreesJack H. Good, Torin Kovach, Kyle Miller, Artur DubrawskiNeurIPS 2023 · 16 citations
