Research
Learning Sparse Decision Trees via Transformer Variational Auto-Encoders
Overview Research area: Interpretable machine learning, specifically decision tree (DT) learning, combined with generative modeling (transformer-based variational auto-encoders) and latent space optim

- arXiv
- 2609.01430
- Published
- 2026-09-01
- Authors
- Giacomo Fidone, Alessio Cascione, Riccardo Guidotti
AI summary
Overview
- Research area: Interpretable machine learning, specifically decision tree (DT) learning, combined with generative modeling (transformer-based variational auto-encoders) and latent space optimization for structured data.
- Technical level: Advanced. The paper assumes familiarity with VAEs, the reparameterization trick, ELBO objectives, transformer attention mechanisms, and the standard trade-offs of decision tree induction algorithms.
- Scope: The paper introduces TREVIS, a framework that learns a continuous latent space of decision trees with a Tree Transformer Variational Auto-Encoder and then searches that space with gradient ascent to produce trees that jointly optimize predictive performance and structural sparsity.
What This Paper Is About
Learning an optimal decision tree is intractable because the discrete space of possible trees grows exponentially with the number of features and training instances, so practical algorithms such as CART, ID3 and C4.5 use greedy top-down splitting and usually return suboptimal trees, while globally optimal solvers cost too much to run. Existing learners also concentrate almost exclusively on predictive performance, with little or no joint control over other desirable properties such as structural sparsity, which governs how interpretable a tree actually is. TREVIS addresses both problems by embedding the discrete space of trees into a continuous latent space learned by a transformer-based VAE, then optimizing that latent space with gradients to find trees that balance performance against a sparsity penalty.
Key Contributions
- TREVIS, a latent-space framework for DT learning. A method that learns tree representations from variational inference in latent space, replacing discrete combinatorial search over trees with continuous optimization over latent vectors.
- TTVAE, a Tree Transformer Variational Auto-Encoder. A VAE in which both encoder and decoder are transformers operating on depth-first pre-order tokenizations of decision trees, using absolute tree positional embeddings to inject hierarchical structure, and injecting the sampled latent vector into every decoder layer through multi-head cross-attention.
- A gradient-based surrogate optimization scheme. Rather than sample-inefficient black-box search over the latent space, TREVIS trains a differentiable surrogate model (an MLP) to approximate the objective value of decoded trees, enabling gradient ascent to steer latent candidates toward high-scoring regions.
- An empirical demonstration on the performance-sparsity trade-off. Experiments on 18 benchmark datasets showing that TREVIS discovers trees whose predictive performance matches near-optimal learners while achieving the lowest average structural complexity among the compared methods.
Main Findings
-
TREVIS achieves competitive predictive performance. On weighted F1-score averaged over the 18 datasets, TREVIS d (the variant trained on discretized features) reaches .841, tied with GOSDT LB (.841), above DL8.5 LB (.834) and FLOW (.791), and below DL8.5 (.851). TREVIS c (trained on the original continuous features) averages .807, below CART c (.832) but above FLOW.
-
TREVIS produces the sparsest trees. Both variants yield the lowest average number of leaves: TREVIS c at 5.72 and TREVIS d at 8.94, compared to CART c (11.61), CART d (11.89), GOSDT LB (11.06), FLOW (10.11), DL8.5 LB (17.44) and DL8.5 (19.39).
-
TREVIS d offers the best overall trade-off. The paper reports that TREVIS d lands in the upper-left region of the performance-versus-leaves plot and provides the most favorable balance of the two objectives. Per the critical difference plots, TREVIS d ranks better than CART d in predictive performance, is statistically indistinguishable from the best near-optimal learners, and attains the best rank for structural sparsity (Nemenyi test at α = 0.1).
-
The surrogate model is a reliable stand-in for the true objective. For TREVIS d, across all datasets, the MLP surrogate achieves an average test Pearson correlation of 0.753 ± 0.146 and an average test RMSE of 0.027 ± 0.018. Across TREVIS overall, the averages are a Pearson correlation of 0.675 ± 0.157 and an RMSE of 0.041 ± 0.031.
-
Runtime is competitive with near-optimal learners but far above greedy ones. Average total runtime is 622.619 seconds for TREVIS c and 490.857 seconds for TREVIS d, against 0.025 s for CART c, 0.013 s for CART d, 207.090 s for DL8.5, 259.747 s for DL8.5 LB, 976.346 s for GOSDT LB and 3437.446 s for FLOW. Within TREVIS c, the breakdown is 360.286 s for TTVAE training, 63.138 s for surrogate training and 199.195 s for the gradient-based search; for TREVIS d it is 268.883 s, 59.156 s and 162.819 s respectively.
-
Some competitors fail to finish within the one-hour budget. GOSDT LB exceeds the one-hour time limit on the
lrsandmagicdatasets, and FLOW exceeds it on every dataset exceptiris. All other competitors complete within the limit. -
Tree positional embeddings matter. The paper states that the effectiveness of the tree absolute positional embeddings is shown empirically in Section IV-D, and that a training collection of 20,000 trees is sufficient for the TTVAE to achieve strong generation quality.
-
Continuous versus discretized feature spaces differ. TREVIS c underperforms CART c, while TREVIS d matches GOSDT LB. The discretization strategy (taken from prior work) restricts candidate splits to those likely to be informative and reduces the vocabulary size, at the cost of optimality with respect to the original feature space.
-
Latent space analysis is only partially reported. The provided content truncates Section IV-C; the paper references UMAP projections of the latent space for four datasets, colored by DT properties, with red lines marking trajectories obtained by moving along the gradient of the corresponding surrogate. No numeric results for this analysis appear in the available content.
Methodology in Plain English
The approach has four stages.
1. Turn trees into text. Each decision tree is written out as a sequence of tokens using a depth-first pre-order traversal. An internal node becomes two consecutive tokens: one for the feature it splits on and one for the threshold value. A leaf becomes a special <L> token. Because every threshold inside an interval between two consecutive feature values produces the same partition of the data, only one canonical value per interval is kept — the left endpoint — and values are rounded to a fixed floating precision. This keeps the vocabulary small. Thresholds are stored as discrete tokens rather than as real numbers, deliberately avoiding an unnecessarily large search space.
2. Learn a latent space of trees. The sequences feed a VAE whose encoder and decoder are both transformers. Since self-attention is permutation-invariant, the authors add absolute tree positional embeddings that encode the root-to-node path as stacked one-hot chunks, with each level scaled geometrically. The encoder reads the sequence with a prepended <CLS> token and produces the mean and variance of the latent distribution; the decoder reads a shifted sequence framed by <BOS> and <EOS> with causal masking, so generation is autoregressive. A latent vector is sampled with the reparameterization trick and injected into every decoder block via multi-head cross-attention. Training maximizes a β-weighted ELBO, with β held at zero for the first 10 epochs and then linearly increased to prevent the KL vanishing problem, plus free bits to lower-bound each latent dimension's KL contribution.
3. Fit a cheap stand-in for the objective. The quantity TREVIS wants to maximize is the weighted F1-score of a decoded tree on the training set minus λ times its number of leaves. That quantity is not differentiable with respect to the latent vector. Instead of calling the decoder repeatedly (as black-box search would), TREVIS trains a small MLP to predict the objective value from the latent representation.
4. Search with gradients and decode. Fifty thousand latent vectors are sampled from a standard normal prior, each is pushed uphill along the gradient of the surrogate for 10 steps, and the top 500 candidates by surrogate score are decoded into trees and scored with the real objective on the training set. This whole procedure is repeated for six values of λ (0.0, 0.0001, 0.0005, 0.001, 0.005, 0.01), and the final tree is the one with the best weighted F1-score on a held-out validation set.
Why This Matters
The work reframes decision tree learning as search in a learned continuous space rather than in a discrete combinatorial one, and shows that this reframing can buy structural simplicity without giving up accuracy. That matters because the discipline is increasingly asked not just for accurate trees but for trees that satisfy multiple properties at once, and today's algorithms mostly optimize only accuracy.
Real-world applications cited or implied by the paper:
- Credit risk assessment, where transparent decision logic is needed to justify lending decisions.
- Hiring, where models influencing employment decisions must be explainable and accountable.
- Healthcare, where high-stakes clinical decision support requires interpretable rule-based reasoning.
- General tabular data pipelines, where decision trees remain a standard model class and controlling tree size directly controls downstream interpretability.
Industry relevance: regulated sectors that require documented, human-readable decision logic can benefit from models with fewer leaves and comparable accuracy. The paper's runtime profile — competitive with near-optimal solvers and much cheaper than GOSDT LB and FLOW on average — suggests the approach is viable where near-optimal learning is already being used. The authors also frame the framework as extensible to fairness, privacy and robustness objectives, which are common compliance requirements. The work is accepted for publication at ICDM 2026.
Future Directions
- Extending TREVIS beyond performance and sparsity. The authors explicitly leave to future work the extension of the framework to more complex objectives, including fairness with respect to protected attributes, privacy by masking sensitive information from adversarial attacks, and robustness to noisy input examples.
- Improving latent space exploration and the surrogate. The surrogate reported average test Pearson correlations of 0.753 ± 0.146 (TREVIS d) and 0.675 ± 0.157 (overall), leaving room for better differentiable approximations and more efficient navigation.
- Closing the gap between continuous and discretized settings. TREVIS c underperformed CART c on average while TREVIS d matched GOSDT LB, which raises the question of how to tokenize thresholds and features so that the continuous-feature variant does not lose accuracy.
- Scaling and efficiency. With six λ runs, 50,000 sampled latents, 10 gradient steps each, and decoding of the top 500 candidates per dataset, the pipeline is substantially slower than CART; reducing this cost or improving competitor parity on large datasets is an open practical question.
- Completing and extending the latent space analysis. The provided content truncates Section IV-C after the description of UMAP projections, so the full quantitative characterization of latent directions and trajectory behavior remains to be fully assessed.
Target Audience
Researchers and practitioners in interpretable machine learning, especially those working on decision tree induction, near-optimal tree solvers, and the performance-interpretability trade-off. It will also interest researchers applying VAEs, transformers and latent space optimization to structured or discrete objects such as graphs and trees, and machine learning engineers in regulated domains who need compact, auditable tree models without sacrificing accuracy. Readers should be comfortable with variational inference, attention mechanisms, and standard tree-induction baselines to follow the methodology and experimental comparisons.
Authors’ abstract
Decision trees are among the most widely used models in machine learning, largely due to their transparent decision logic, making them well-suited for high-stakes decision-making contexts. However, most existing learning algorithms focus on predictive performance, overlooking the joint optimization of other desirable properties, such as structural sparsity. In this work we propose TREVIS, an approach for learning decision trees with respect to complex objectives, based on the exploration of the latent space of a Tree Transformer Variational Auto-Encoder (TTVAE). By mapping decision trees onto latent representations, TREVIS replaces the discrete search space with a continuous one, enabling gradient-based optimization via a differentiable surrogate model. We experiment with TREVIS for learning decision trees that jointly optimize predictive performance and sparsity. Results show that TREVIS discovers decision trees matching the predictive performance of existing near-optimal algorithms while improving their structural sparsity.