Deep learning on tabular data
Be able to decide when neural networks beat tree models on tabular data — and when they do not.
Prerequisites
Intuition
Deep learning dominates images, text, audio and video. On tabular data the situation is a different one: gradient boosting (XGBoost, LightGBM, CatBoost) still wins most comparisons.
That is a well-documented result, not an opinion. Grinsztajn et al. (2022) compared systematically across 45 datasets and found that tree ensembles beat neural networks in most cases — even after careful hyperparameter search for both.
Three explanations:
- Tabular data has irregular decision boundaries. Trees split on thresholds and can make sharp jumps; neural networks have a built-in preference for smooth functions.
- Irrelevant columns. Real tabular data contains a lot of noise. A tree ignores a useless column entirely; a neural network still gets it as input.
- No spatial or sequential structure to exploit. The order of the columns is arbitrary, so there is nothing for an architecture to build on — and hence no inductive bias to gain from.
Formal
When a neural network is the right choice anyway:
| The situation | Why |
|---|---|
| Very large data (millions of rows) | the network has time to learn; the trees' advantage shrinks |
| Many categories | learnt embeddings beat one-hot |
| Mixed modalities | a table + text + an image in one model |
| Multi-task learning | shared representations between related targets |
| Transfer learning | pre-train on one domain, fine-tune on another |
| Part of a larger network | when the table is the input to something else |
The third row is the most common real reason: as soon as a free-text column is to be used together with the numerical ones, a neural network becomes natural.
A fair comparison requires four things, and they are missing from most blog posts:
| The requirement | Why |
|---|---|
| The same hyperparameter budget for both | the network needs more searching to reach its potential |
| The same preprocessing where it is relevant | but trees need no scaling |
| Several datasets | one dataset says nothing |
| Several seeds, with the spread reported | the differences are often smaller than the noise |
Architectures built for tabular data: TabNet, FT-Transformer, SAINT and NODE. They are interesting and can beat boosting on individual datasets — but in broad comparisons the result is usually equivalent or worse, at a considerably higher cost in time and complexity.
The practical recommendation is simple:
- Start with gradient boosting. It is fast, robust against untidy data and needs little tuning.
- If you have one of the situations in the table above — try a network with embeddings.
- Ensemble the two if you need the last per cent. Their errors are often different, and the average beats both.
And what beats both: better features. On tabular data, domain knowledge is still the most profitable investment — a ratio or a time difference suggested by somebody who knows the field more often gives more than any change of model.
Code
import numpy as np, torch, torch.nn as nn
from sklearn.model_selection import cross_val_score, KFold
from sklearn.ensemble import HistGradientBoostingClassifier
from sklearn.linear_model import LogisticRegression
from sklearn.pipeline import Pipeline
from sklearn.preprocessing import StandardScaler
def compare_fairly(X, y, seeds=(0, 1, 2)):
"""The same cross-validation, several seeds, the spread reported."""
models = {
"boosting": lambda s: HistGradientBoostingClassifier(random_state=s),
"linear": lambda s: Pipeline([("s", StandardScaler()),
("m", LogisticRegression(max_iter=2000))]),
}
for name, build in models.items():
scores = []
for s in seeds:
cv = KFold(5, shuffle=True, random_state=s)
scores += list(cross_val_score(build(s), X, y, cv=cv, scoring="roc_auc"))
print(f" {name:<9} {np.mean(scores):.4f} ± {np.std(scores):.4f} (n={len(scores)})")
# A neural network for tabular data: embeddings for the categories, normalisation for the numbers
class TabularNet(nn.Module):
def __init__(self, categories: dict[str, int], n_num: int, hidden=256):
super().__init__()
def dim(n):
return min(600, int(round(1.6 * n ** 0.56)))
self.emb = nn.ModuleDict({k: nn.Embedding(n + 1, dim(n)) for k, n in categories.items()})
self.num_norm = nn.BatchNorm1d(n_num)
width = sum(e.embedding_dim for e in self.emb.values()) + n_num
self.f = nn.Sequential(
nn.Linear(width, hidden), nn.ReLU(), nn.Dropout(0.3),
nn.Linear(hidden, hidden // 2), nn.ReLU(), nn.Dropout(0.3),
nn.Linear(hidden // 2, 1),
)
def forward(self, cat, num):
parts = [self.emb[k](v) for k, v in cat.items()] + [self.num_norm(num)]
return self.f(torch.cat(parts, dim=1)).squeeze(-1)
# Ensemble: the errors are often different, so the average beats both
def ensemble(p_boosting, p_net, weight=0.5):
return weight * np.asarray(p_boosting) + (1 - weight) * np.asarray(p_net)
# Check that the errors really are different — otherwise the ensemble is pointless
def error_correlation(y, p1, p2):
return float(np.corrcoef(np.abs(y - p1), np.abs(y - p2))[0, 1])
# < 0.7 means the models make different errors and that the ensemble should help
The last function is worth running before you build an ensemble: if the error correlation is 0.95 the models make the same mistakes, and putting them together gives nothing.
Mastery means
- Compares neural networks and tree models fairly
- Knows why trees usually win on tabular data
- Knows when a neural network is the right choice anyway
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Why do tree-based models still outperform deep learning on typical tabular data? — arXiv (open access; licence per article)
- scikit-learn User Guide (BSD-3) — BSD-3-Clause
- arXiv — Revisiting Deep Learning Models for Tabular Data — arXiv (open access; licence per article)