Why do tree-based models still outperform deep learning on tabular data? (2022)
arxiv.org
arxiv.org
There is no analog for tabular data, it's all different.
Similarly, tabular data is often of this nature. Its not i.i.d, it tends to cluster.
>Creating tabular-specific deep learning architectures is a very active area of research (see section 2) given that tree-based models are not differentiable, and thus cannot be easily composed and jointly trained with other deep learning blocks.
Here is a second reason, from the paper
>Impressed by the superiority of tree-based models on tabular data, we strive to understand which inductive biases make them well-suited for these data.
which is a great reason, because understanding the inductive biases of different learning/regression techniques gets us closer to a more general understanding of how to encode inductive biases in a generic learning algorithm.
In the domains where NNs work well (image processing and language), you're dealing with a predictable and stable distribution of values. Elephants might look a bit different in the train and test set, but you're not randomly getting 100x the variance of the input data. The decision tree just isn't going to care as much, because splits around the mean will lead to the same outcome.
Another hypothesis is that zooming into bivariable relationships is more important in tabular data. Neural nets are better at local and global context. But they struggle if all that matters is the relationship between two columns of data because of the additive nature. Large networks can figure it out due to model capacity, but then you'll run into overfitting.
1. Something like a deep support vector machine. Instead of (linear) -> (any activation), you want to create a bunch of features that look like testing the vector against a splitting hyperplane. One option is (bias) -> (matmul) -> (1-bit sigmoid). Applying a bias term _for each row_ let's you choose the branch location, the matmul's result will be positive or negative at each output feature depending on which side of the hyperplane normal to the vector described by the corresponding row you happen to fall on. Then just bring that down to -1 or 1 so you can't sneak much nonstationary drift variance into the output (perhaps train with a normal sigmoid annealed to behave more like this one, and a suitable regularizing term to keep the network from sneaking in values near 0 to thwart your annealing).
2. Use an attention-like mechanism, but across features (this would likely require an additional tensor channel, so that each "feature" carries information in a high enough dimensional space for this to do something meaningful). You apply the inductive bias that sparse feature interactions are important and need to be discovered.
Those two ideas also compose easily.
Suppose input data is [batch_size, num_features]. Then you do x.unsqueeze(1) giving you [batch_size, num_features, 1]. Then what?
einsum('bf,fc->bfc', batched_inputs, channel_embedding)
Then carry that info through the network and project it down at the end. It's roughly equivalent to the token embedding step in an LLM.
E.g. finance
In a sufficiently competitive space, good enough doesn't cut it.
They are also more attractive for streaming data. Tree-based models can't learn incrementally. They have to be retrained from scratch each time.
You might call this overfitting/noise/.... but if you do it carefully it's profitable.
- Explainability / debug-ability of models
- Effort to train, deploy, and manage NN models in production
- Capturing, collating, and organizing new & better datasets
- Local developer experience and human-model-iteration time
Building all of your software in C or Assembly will be faster and higher performant. But at what cost and with what tradeoffs? Building a website has a different set of tradeoffs than building a program for the Mars rover.
Or, if the "tabular data" is heavily relationship-based, then possibly replace "relational data warehouse" with "graph database", and "SQL queries" with whatever querying language that graph DB is natively / most expressively queried in.
Of course, this is the most important implicit "equally important" factor, one that an ML dev would think goes without mentioning: the generality or "power" of the model in what questions it can answer. You can only make these trade-offs in the context of knowing what kinds of questions you want your model to solve for! If all your questions are quantitative ones, maybe the right "model" for you is an RDBMS!
---
Though, that being said... why can't a deep-learning model emulate the thing that an RDBMS does, "at runtime", as part of its "mental toolkit" for approaching problems? That would be the best of both worlds, no?
I know that LLMs in particular have been observed to have "emergent numeracy" above a certain training-set size. There is a step function in how they approach such problems, going from their only being able to answer arithmetic questions on numbers of bounded size, and sometimes getting the answers wrong (probably this is due to a memorization-based approach); to being able to answer arbitrary arithmetic questions on operands of unbounded size, and always getting the answer correct.
I would guess that that what's happening, is that they are developing a functional component of their network that works akin to an Arithmetic Logic Unit, operating not on tokens, but on tokens transformed into a "numeric register" representation that is amenable to having math done to it with stable, quantized, position-independent results. (Just like the functional component that human brains develop after seeing enough math problems... probably.)
Do you, as an ML dev, think it would ever be possible for any of the model architectures we're familiar with today, to be trained such that they would develop an analogous emergent functional component for handling tabular-data questions, by transforming its internal working state into relational-DB/graph-DB data structures — e.g. page-heaps of binary-packed row-tuples; B-tree indices; etc — and then manipulating the working state in that form, using learned algorithms applicable to that type of data?
It seems to me (possibly just because I don't know any better) that just as with numeracy, "being able to put the data into a different and better internal representation" is what would be needed for deep-learning models to become truly good at dealing with tabular-data problems.
But, unlike with numeracy, "thinking as if you were a relational database" is not something a single human would ever intuit how to do without being taught. Relational algebra — and the data-structures and algorithms to make it practical to have a Turing machine do said relational algebra — wasn't even a single intuition, but a conscious effort, of multiple humans, working together over years. I strongly doubt that there's any number of "tabular-data problems" that you could show a human being, that would result in them developing an intuitional ability to do what a relational database does with its memory to efficiently answer queries.
(I suppose we could give an ML model an RDBMS, and hardwire it to interact with it. I know there are hybrid ML + formal-logic systems. Are there hybrid ML + data-warehouse systems? Not where the model queries an external DB — while that can be done, it'd be only in the same "stop and do this" way that ChatGPT runs Python code, which wouldn't make it a thinking tool the way that the formal-logic proof engines are for hybrid ML systems. Rather, I mean that some data-warehouse execution engine could be embedded into the ML execution framework itself, deployed as part of the GPU shader-program to each tensor core, such that data-warehouse operations can be done as a native part of the network's per-node instruction-set. Anyone ever tried this?)
It's doubly funny; as someone that comes from an ML background, and has developed and maintained multiple ML systems at multiple orgs, that I also think the answer very often is, "throw the tabular data into a relational data warehouse, and ask your questions in the form of SQL queries."
You can ask SQL descriptive questions. Can you ask it for predictions? How?
One of several examples of implementing linear regression in SQL.
You're correct, but "in some cases" is doing a lot of work here.
With the tooling where it's at, how much harder is it to apply xGBoost vs a linear model?
How do you know which questions to ask? This is what ML is good at, finding the right questions which classify the data.
Maybe we already know everything about the dataset. For example, if it's line-of-business customer data gradually built up by a sales team, then the brains of the salespeople have likely already done all the "implicit classification" needed to generate good questions about the dataset.
And this is, by far, the usual scenario for Business Intelligence questions: someone with "business-domain knowledge", e.g. an executive, has formed an intuitional hypothesis about the data based on their personal experience; and so they ask someone with "data-domain knowledge", e.g. a business analyst or data scientist, to test that hypothesis.
It's actually rare, in my experience, to have a tabular-data dataset that someone is motivated to understand, that doesn't also "come with" a set of people who can already act as (good!) models trained on that dataset, to aid them in that understanding. (Sometimes these people can't find each-other — but they do usually exist.)
AFAIK, having reams of entirely opaque and ill-understood tabular data, such that you need classification/clustering to get started on asking questions, only really happens in the sciences: sensor-network climate data; longitudinal-study medical-outcome data; census data; housing-market data; etc. In other words, it's almost always universities and governments — not businesses — that care about analyzing opaque tabular data.
And that's a key to understanding the constraints in play for choosing models! Because business-driven analyses are usually time-constrained in some way (potentially even needing post-training question-answers to be generated in soft-realtime); while institutional analyses usually aren't. Big difference!
- identify latent features of customers via their behavioral data, to be used for profiling customers or recommending products to them
- within a large amount of customer behavioral data, identify potentially fraudulent behavior
- identify causes of seasonality (e.g. temporal patterns) in the data in order to improve forecasting (sales, traffic, whatever)
In those cases part of the investigation is to initially take a hands-off (unsupervised) approach, so that we can compare our initial top-down hypotheses with actual patterns in the data.
In both of those cases there's considerable (and sometimes adversarial) noise in the data.
Having understood that question, and built an understanding of what predicts fraud, you would then graduate to build models to understand the extent to which features predict fraudulence.
My point in context of the conversation is that it's useful in a business context to explore and understand that data.
https://www.forbes.com/sites/kashmirhill/2012/02/16/how-targ...
It's not clear what your point is. If you're not interested in the predictions that tree-based models provide, do not use tree-based models on your tabular data. A predictive model and a SQL query are not the same thing.
Lately, that means they're often spending a lot of resources (and even novel R&D time!) getting various kinds of ML models trained on the data.
My point is that this is often pointless, because, given the type of data they're working with (tabular, quantitative line-of-business data), they won't actually see "arbitrary questions"; they'll see the strict subset of arbitrary questions that could have been solved just as well — if not much better! — with a SQL query. And for much less capital expenditure — because the LOB data usually already lives in an RDBMS in the first place.
For those curious about what we have been up to on the topic of tabular learning, we have found a setting where deep learning does seem to bring sizable benefits (spoiler alert, it's about being able to pre-train, and transferring to new data works best when there are some strings to be recognized): https://arxiv.org/abs/2402.16785
In the above work, pre-trained tabular models markedly outperform tree-based models (including catboost, which is a very strong baseline).
As someone who has been banging on tabular data for years, I'm really excited about this development.
When tabular data is mentioned, one of the unspoken applications is finance. There, my guess is that one of the issues is that data is not very IID and thus latent "events" are fairly sparse. Combine that with the humongous amount of raw data, and you get models that overfit.
What do you base this on? Having only neural nets on tabular data is mostly done due to laziness of the creator since neural nets are much easier to use, not because neural nets perform better even with large amounts of data. In general you want both since they are good at finding different kinds of patterns.
Consider: Predicting how much a customer might pay by end of month, with information we have at the start of the month.
In this example, if a customer had a record $10m of open invoices due by EoM and the largest payment amount received in prior months of $5m, the decision tree cannot possibly predict the payment amount will be ~$10m, even when the best feature indicates the payment will be $10m.
There are some hacks/techniques which can maybe reduce this issue, but they don't always work.
Also all models are a “mean of the subgroup of the data.” The prediction is by definition the conditional mean as a function of the input values.
PS: Its weird that you are being down-voted. I think your opinion is reasonable.
Retraining - online training solves this for the most part.
Frameworks - the only battle-tested batteries-included one I've seen is Vespa. Noone else publishes any of interesting bits. KDD is the most relevant conference if you're interested in the field. IIRC Xiaohongshu has some papers that can only really be done with NNs.
So I would be curious to see latest DL results. On the other hand it is also the case that in most cases where DL based on foundation models is used, specific heavily tuned models outperform the generalistic models. And for tabular data there is a lot of experience how to make it great with tree based models.
https://arxiv.org/abs/2403.01841 (ICLR 2024 spotlight)
Tree-based models are extremely good at finding clustering patterns; they outperform trained humans at that, thus we have commercial applications such as fraud detection.
Deep Learning is most promising way of getting us to general intelligence. So far only known general intelligence, human intelligence, has many quirks at specific tasks and I think Deep Learning won't be any different. However Deep Learning models can recognise their own weakness and call tree-based model if they think that's appropriate.
(I'm personally using NN models for predicting certain values for tabularly structured data and at least for my case, the NN works better than state-of-the art tree models.)
For example, take the circle dataset here: https://playground.tensorflow.org
That doesn't look immediately linearly separable, but since it is 2D we have the insight that parameterizing by radius would do the trick. Now try doing that in 1000 dimensions. Sometimes you can, sometimes you can't or don't want to bother.
The magic of deep neural networks comes from modeling complicated conditional probability distributions, which lets you do generative magic but isn't going to give you significantly better results than ensemble kNN when you're discriminating and the conditional distribution is low variance. Ensemble methods are like a form of regularization and they also act as a weak bootstrap to better model population variance, so it's no surprise that when they're capable of modeling the domain, they perform better than unregularized, un-bootstrapped neural network model. There are still tons of situations where ensemble methods can't model the domain, and if you incorporated regularization and bootstrapping into a discriminative NN model it would probably perform equivalently to the ensemble model.
Transformers with positional encoding have embeddings are invariant to the input order. CNN's have translation invariance and can have little rotational invariance.
It's harder to find similar invariances to tabular data. Maybe applying methods from GNN's would help?
Ensemble models work well because they reduce both bias & variance errors. Like DTs, NNs have low bias errors and high variance errors when used individually. The variance error drops as you use more learners (DTs/NNs) in the ensemble. Also, the more diverse the learners, the lower the overall error.
Simple ways to promote the diversity of the NNs in the ensemble is to start their weights from different random seeds and train each one of them on a random sample from the overall training set (say 70-80% w/o replacement).
smthin smthin inductive bias?
But they don't scale with larger and more complex data. You cannot (realistically) make an LLM with XGBoost.
Kind of surprised how well Resnet and FT Transformer do though.
https://arxiv.org/abs/1806.06988
It combines NN’s with decision trees.
feel free to extend logistic regression to an MLP :)