Hoppa till innehållet
AI-grafen
E· Universitetdeep-learning· ca 60 min· utvecklande· verifierad 2026-09-20

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:

  1. 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.
  2. Ovidkommande kolumner. Verklig tabelldata innehåller mycket brus. Träd ignorerar en oanvändbar kolumn helt; neuronnät får den ändå som indata.
  3. 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ägeVarför
Mycket stor data (miljoner rader)nätet hinner lära sig; trädens fördel krymper
Många kategorierinlärda embeddingar slår one-hot
Blandade modalitetertabell + text + bild i en modell
Fleruppgiftsinlärningdelade representationer mellan relaterade mål
Överföringsinlärningförträna på en domän, finjustera på en annan
Del av ett större nätnä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:

KravVarför
Samma hyperparameterbudget för bådanätet kräver mer sökning för att nå sin potential
Samma förbehandling där det är relevantmen träd behöver ingen skalning
Flera datasetett dataset säger ingenting
Flera frön, med spridning rapporteradskillnaderna ä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:

  1. Börja med gradient boosting. Det är snabbt, robust mot ostädad data och kräver lite tuning.
  2. Har du en av situationerna i tabellen ovan — pröva ett nät med embeddingar.
  3. 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

Alla källor och licenser