Deep learning på tabelldata
Kunna avgöra när neuronnät slår trädmodeller på tabelldata — och när de inte gör det.
Förkunskaper
Intuition
Djupinlärning dominerar bild, text, ljud och video. På tabelldata är läget ett annat: gradient boosting (XGBoost, LightGBM, CatBoost) vinner fortfarande de flesta jämförelser.
Det är ett välbelagt resultat, inte en åsikt. Grinsztajn m.fl. (2022) jämförde systematiskt över 45 dataset och fann att trädensembler slog neuronnät i de flesta fall — även efter noggrann hyperparametersökning för båda.
Tre förklaringar:
- Tabelldata har oregelbundna beslutsgränser. Träd delar på tröskelvärden och kan göra skarpa hopp; neuronnät har en inbyggd förkärlek för mjuka funktioner.
- Ovidkommande kolumner. Verklig tabelldata innehåller mycket brus. Träd ignorerar en oanvändbar kolumn helt; neuronnät får den ändå som indata.
- Ingen rumslig eller sekventiell struktur att utnyttja. Kolumnernas ordning är godtycklig, så det finns inget för en arkitektur att bygga på — och därmed ingen induktiv bias att vinna på.
Formellt
När neuronnät ändå är rätt val:
| Läge | Varför |
|---|---|
| Mycket stor data (miljoner rader) | nätet hinner lära sig; trädens fördel krymper |
| Många kategorier | inlärda embeddingar slår one-hot |
| Blandade modaliteter | tabell + text + bild i en modell |
| Fleruppgiftsinlärning | delade representationer mellan relaterade mål |
| Överföringsinlärning | förträna på en domän, finjustera på en annan |
| Del av ett större nät | när tabellen är indata till något annat |
Den tredje raden är den vanligaste verkliga anledningen: så snart en fritextkolumn ska användas tillsammans med de numeriska blir ett neuralt nät naturligt.
En rättvis jämförelse kräver fyra saker, och de saknas i de flesta blogginlägg:
| Krav | Varför |
|---|---|
| Samma hyperparameterbudget för båda | nätet kräver mer sökning för att nå sin potential |
| Samma förbehandling där det är relevant | men träd behöver ingen skalning |
| Flera dataset | ett dataset säger ingenting |
| Flera frön, med spridning rapporterad | skillnaderna är ofta mindre än bruset |
Arkitekturer byggda för tabelldata: TabNet, FT-Transformer, SAINT och NODE. De är intressanta och kan slå boosting på enskilda dataset — men i breda jämförelser är resultatet oftast likvärdigt eller sämre, till betydligt högre kostnad i tid och komplexitet.
Den praktiska rekommendationen är enkel:
- Börja med gradient boosting. Det är snabbt, robust mot ostädad data och kräver lite tuning.
- Har du en av situationerna i tabellen ovan — pröva ett nät med embeddingar.
- Ensembla de två om du behöver sista procenten. Deras fel är ofta olika, och medelvärdet slår båda.
Och det som slår båda: bättre features. På tabelldata är domänkunskap fortfarande den mest lönsamma investeringen — ett förhållande eller en tidsskillnad som någon som kan området föreslår ger oftare mer än något modellbyte.
Kod
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 jamfor_rattvist(X, y, fron=(0, 1, 2)):
"""Samma korsvalidering, flera frön, spridning rapporterad."""
modeller = {
"boosting": lambda f: HistGradientBoostingClassifier(random_state=f),
"linjär": lambda f: Pipeline([("s", StandardScaler()),
("m", LogisticRegression(max_iter=2000))]),
}
for namn, bygg in modeller.items():
poang = []
for f in fron:
cv = KFold(5, shuffle=True, random_state=f)
poang += list(cross_val_score(bygg(f), X, y, cv=cv, scoring="roc_auc"))
print(f" {namn:<9} {np.mean(poang):.4f} ± {np.std(poang):.4f} (n={len(poang)})")
# Neuronnät för tabelldata: embeddingar för kategorier, normalisering för numeriska
class TabellNat(nn.Module):
def __init__(self, kategorier: dict[str, int], n_num: int, dolt=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 kategorier.items()})
self.num_norm = nn.BatchNorm1d(n_num)
bredd = sum(e.embedding_dim for e in self.emb.values()) + n_num
self.f = nn.Sequential(
nn.Linear(bredd, dolt), nn.ReLU(), nn.Dropout(0.3),
nn.Linear(dolt, dolt // 2), nn.ReLU(), nn.Dropout(0.3),
nn.Linear(dolt // 2, 1),
)
def forward(self, kat, num):
delar = [self.emb[k](v) for k, v in kat.items()] + [self.num_norm(num)]
return self.f(torch.cat(delar, dim=1)).squeeze(-1)
# Ensembla: felen är ofta olika, så medelvärdet slår båda
def ensemble(p_boosting, p_nat, vikt=0.5):
return vikt * np.asarray(p_boosting) + (1 - vikt) * np.asarray(p_nat)
# Kontrollera att felen faktiskt är olika — annars är ensemblen meningslös
def korrelation_mellan_fel(y, p1, p2):
return float(np.corrcoef(np.abs(y - p1), np.abs(y - p2))[0, 1])
# < 0.7 betyder att modellerna gör olika fel och att ensemblen bör hjälpa
Sista funktionen är värd att köra innan man bygger en ensemble: är felkorrelationen 0,95 gör modellerna samma misstag, och att slå ihop dem ger ingenting.
Behärskning innebär
- Jämför neuronnät och trädmodeller rättvist
- Känner till varför träd oftast vinner på tabelldata
- Vet när neuronnät ändå är rätt val
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Why do tree-based models still outperform deep learning on typical tabular data? — arXiv (öppen åtkomst; licens per artikel)
- scikit-learn User Guide (BSD-3) — BSD-3-Clause
- arXiv — Revisiting Deep Learning Models for Tabular Data — arXiv (öppen åtkomst; licens per artikel)