Performance#
How long a soft decision tree takes to fit, measured rather than estimated.
One CPU thread (OMP_NUM_THREADS=1), Apple silicon, synthetic data with
50 features and 5 classes, depth 4, batch size 64, 30 epochs.
samples |
fit, 30 epochs |
per epoch |
predict (all samples) |
|---|---|---|---|
1 000 |
2.0 s |
66 ms |
3 ms |
5 000 |
6.3 s |
210 ms |
13 ms |
20 000 |
30 s |
1.0 s |
39 ms |
50 000 |
58 s |
1.9 s |
102 ms |
Time is linear in the number of samples and, since every gate is evaluated for every sample, linear in the number of gates: depth 6 (63 gates) costs about 1.4 times depth 4 (15 gates) at 5 000 samples. At the default 150 epochs a depth-4 tree on 50 000 samples fits in about five minutes on one thread.
Where the time goes#
About a third of a training step is the backward pass and a fifth was, until
0.7, the DataLoader indexing the dataset one sample at a time. Batches are
now sliced directly from the tensors (neural_trees._batching), which
draws the same permutation the loader would draw from the same seed, so fits
with a random_state are bit-identical to the previous version’s and
14–21% faster.
What is not fast#
Growth.
growth="incremental"trains up todepthrounds andgrowth="per_leaf"up to2 * depth; withgrowth_budget="full"each round gets the whole epoch budget, so the fit costs that many times a fixed tree.Omnivariate trees cross-validate an MLP at every node. They are for small problems.
Batch size. Larger batches are faster per epoch and reach a worse model at the same number of epochs (on Digits, 20 epochs: batch 64 gives 0.915 training accuracy in 1.6 s, batch 1024 gives 0.786 in 0.5 s). Raise
max_epochswith the batch size.
Devices#
device="auto" uses CUDA, then Apple MPS, then CPU. On problems of the
sizes above the CPU is not the bottleneck and a GPU rarely helps; it starts
to pay at hundreds of thousands of samples or deep trees.
Prediction#
How fast a fitted tree predicts, by export, one thread. One row is the
latency a service answering single requests pays; batch is 10 000 rows
at once. The torch estimator, the numpy copy (to_numpy()), ONNX Runtime
on the ONNX export (to_onnx()) and the hard rule tree (to_hard_tree())
are the same fitted tree; the first three predict the same probabilities to
float32 precision, the last reads each gate as a hard decision. Measured by
benchmarks/latency.py, best of seven, Apple silicon.
Export |
Breast Cancer, depth 4, 30 features |
Digits, depth 6, 64 features |
||
|---|---|---|---|---|
one row |
batch, rows/s |
one row |
batch, rows/s |
|
torch estimator |
93 µs |
2,485,509 |
111 µs |
448,645 |
numpy export |
33 µs |
1,176,864 |
50 µs |
205,547 |
ONNX Runtime |
11 µs |
3,022,593 |
33 µs |
565,836 |
hard rule tree |
44 µs |
7,259,967 |
58 µs |
2,974,960 |
Single-row latency is dominated by call overhead, not arithmetic: a depth-4 tree is fifteen dot products. ONNX Runtime has the least overhead and is the export to serve one request at a time; the hard rule tree is the fastest in batch because it evaluates one path instead of every gate. Nothing here needs a GPU, and the ONNX file needs neither torch nor this library where it runs.