Accueil / Factory / Algos ML / TabNet — factory / algos ML / deep learning sur données tabulaires

RÉSEAU TABNET.

TabNet est un réseau de neurones conçu par Google pour les données en tableau. À chaque étape de décision, un mécanisme d'attention choisit les quelques variables à regarder, client par client. Il donne ainsi une importance des variables par ligne, mais bat rarement un bon boosting sur ce terrain.

Données tabulairesDeep learningClassificationAttentionNiveau : avancé

FICHE D'IDENTITÉ

notes sur 5 · usage entreprise
PerformanceProche du boosting au mieux, souvent un cran en dessous
InterprétabilitéMasques d'attention par client, à lire avec recul
VitesseEntraînement bien plus long qu'un XGBoost, GPU utile
Facilité de réglageNombreux hyperparamètres, sensible à leur réglage
Tolérance aux données brutesVariables catégorielles gérées ; valeurs manquantes à traiter
EN 30 SECONDES

Un conseiller étudie un dossier client en plusieurs passes. À la première, il ne regarde que le contrat et les incidents. À la suivante, il vérifie l'ancienneté. Il note à chaque passe ce qu'il a regardé.

1. Une étape choisit ses variables

Un module d'attention produit un masque : un poids par variable, la plupart exactement à zéro (fonction sparsemax). Seules les variables retenues passent à la suite, et ce choix change d'un client à l'autre.

2. Les étapes s'enchaînent

Chaque étape transforme les variables retenues et apporte sa part à la décision. Un coefficient limite la réutilisation des mêmes variables d'une étape à l'autre, pour pousser le modèle à en explorer d'autres.

3. On additionne et on lit les masques

Les contributions des étapes s'additionnent pour donner la probabilité finale. Les masques cumulés donnent l'importance de chaque variable, pour un client ou pour toute la base.

LE CAS MÉTIER

churn · télécom / énergie / assurance
EN ENTRÉE

Une ligne par client

Ancienneté, appels au support, montant mensuel, type de contrat, incidents des 3 derniers mois, et la colonne à prédire : le client est-il parti ?

EN SORTIE

Un score et les variables regardées

Chaque client reçoit une probabilité de départ et un masque qui indique sur quelles variables le modèle s'est appuyé pour lui. L'importance globale résume ces masques sur toute la base.

CE QU'ON MESURE

L'AUC, comparée à celle d'un boosting

L'AUC mesure si les partants reçoivent un score plus haut que les fidèles (0,5 : hasard, 1 : perfection). Le vrai test est la comparaison avec un XGBoost entraîné sur le même découpage : si TabNet ne fait pas mieux, sa complexité ne se justifie pas.

QUAND LE SORTIR, QUAND L'ÉVITER

OUI

  • Très gros volumes tabulaires (des millions de lignes) où le deep learning a de la matière
  • Besoin d'une importance des variables par client, intégrée au modèle
  • Données mêlant tableau et autres sources, dans une architecture de deep learning commune
  • Pré-entraînement non supervisé possible sur beaucoup de lignes non étiquetées

NON

  • Quelques milliers de lignes : XGBoost, LightGBM ou une Random Forest font presque toujours mieux, plus vite
  • Besoin d'une explication auditable : SHAP sur un modèle d'arbres est plus établi
  • Pas de GPU ni de temps pour régler les hyperparamètres
  • Première modélisation d'un problème : commencer par une régression logistique ou un boosting
LES 4 RÉGLAGES QUI COMPTENT

Noms donnés pour R (tabnet) et Python (pytorch-tabnet).

num_steps / n_steps

Le nombre d'étapes de décision, donc de passes sur les variables. 3 à 7 en général. Plus d'étapes : plus de capacité, plus de risque de surapprentissage.

decision_width, attention_width / n_d, n_a

La largeur des couches de décision et d'attention. On les garde souvent égales, de 8 à 64.

penalty / lambda_sparse

La force qui pousse les masques à ne retenir que peu de variables. Plus elle est forte, plus les masques sont lisibles, parfois au prix de la performance.

batch_size, virtual_batch_size

TabNet utilise une normalisation par sous-lots (ghost batch normalization). Des lots assez grands stabilisent l'entraînement ; le sous-lot, plus petit, fixe la taille des groupes normalisés.

LE CODE MINIMAL

jeu d'exemple : clients_churn.csv ↓
# Churn clients : TabNet en R
library(tabnet)

clients <- read.csv("clients_churn.csv")
clients$churn <- factor(clients$churn)
clients$contrat <- factor(clients$contrat)

set.seed(42)
torch::torch_manual_seed(42)  # graine du moteur torch sur lequel repose tabnet
idx <- sample(nrow(clients), round(0.7 * nrow(clients)))
train <- clients[idx, ]
test <- clients[-idx, ]

modele <- tabnet_fit(churn ~ anciennete + appels_support + montant + contrat + incidents_3m,
                     data = train, epochs = 50, batch_size = 256, virtual_batch_size = 128,
                     valid_split = 0.2, verbose = FALSE)

# Score de départ sur les clients jamais vus
proba <- predict(modele, test, type = "prob")$.pred_1

# AUC calculée en R base (formule des rangs)
rangs <- rank(proba)
n1 <- sum(test$churn == "1")
n0 <- sum(test$churn == "0")
cat("AUC test :", round((sum(rangs[test$churn == "1"]) - n1 * (n1 + 1) / 2) / (n1 * n0), 3), "\n")

# Importance globale des variables, issue des masques d'attention
print(modele$fit$importances)

QUESTIONS FRÉQUENTES

TabNet est-il meilleur que XGBoost ?

Rarement. Plusieurs études comparatives publiées depuis 2021 montrent que les boostings d'arbres restent en général devant les réseaux de neurones, TabNet compris, sur les données tabulaires de taille courante. TabNet peut rivaliser sur de très gros volumes ou quand il est combiné à d'autres modèles.

Comment TabNet explique-t-il ses prédictions ?

À chaque étape, un masque attribue un poids à chaque variable, avec beaucoup de zéros. En cumulant ces masques, on obtient l'importance des variables pour chaque client et pour l'ensemble. C'est une indication de ce que le modèle a regardé, pas une mesure d'effet causal.

Faut-il un GPU pour entraîner TabNet ?

Pas obligatoirement : quelques milliers de lignes s'entraînent sur un processeur en quelques minutes. Sur des millions de lignes, le GPU devient vite indispensable.

LES ALGOS VOISINS

à comparer avant de choisir
la référence à battre

XGBoost

Le boosting d'arbres reste en général le meilleur choix sur des données tabulaires, pour bien moins de réglages et de calcul.

Voir la fiche →
le réseau sans attention

Réseau de neurones (MLP)

Le réseau dense classique appliqué au tableau. Plus simple que TabNet, sans sélection de variables intégrée.

Voir la fiche →
l'explication par client

SHAP

Attribue à chaque variable sa contribution au score d'un client, pour n'importe quel modèle. Plus établi que les masques d'attention.

Voir la fiche →
— formation

Passer de la fiche à la pratique

Dataistudio forme les équipes au machine learning et à l'IA, sur des cas concrets.

Voir les formations →