Soft Classification Trees with Conditional Computation
Abstract
Decision trees are widely used supervised learning models for tabular data. Multivariate deterministic trees provide interpretable decision rules and conditional computation but typically rely on expensive mixed-integer optimization, whereas soft trees are amenable to continuous optimization at the cost of losing single-path routing and conditional computation. We propose a soft classification tree that restores single-path routing and conditional computation while retaining the favorable properties of continuous optimization. Our approach relies on a new activation function for routing samples at branch nodes, defined as a novel scaled variant of the negative log-likelihood that combines its benefits with those of the sigmoid, is differentiable almost everywhere, and is proven to be classification-calibrated. The resulting model is suitably trained using a projected adaptation of the Adam optimizer to handle class assignment constraints. Unlike prior differentiable trees, our method needs only a single first-order run, with no scaling, annealing, warm start, or multi-start. Experiments on 61 classification datasets show that our single-tree method is better than all the compared multivariate single-tree alternatives, attaining essentially the same predictive performance as OCT-H, the strongest multivariate tree baseline, while requiring much lower average training time. It also achieves predictive performance competitive with Random Forest and within about on average of XGBoost. On synthetic datasets, it substantially outperforms OCT-H in recovering the underlying ground-truth tree structure.
est. 32% chance this paper gets accepted at ICLR 2027.
What do you think this paper will get?
All positions stay anonymous.