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.
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é.
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.
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.
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.
Ancienneté, appels au support, montant mensuel, type de contrat, incidents des 3 derniers mois, et la colonne à prédire : le client est-il parti ?
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.
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.
Noms donnés pour R (tabnet) et Python (pytorch-tabnet).
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.
La largeur des couches de décision et d'attention. On les garde souvent égales, de 8 à 64.
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.
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.
# 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)
# Churn clients : TabNet en Python
import pandas as pd
from pytorch_tabnet.tab_model import TabNetClassifier
from sklearn.model_selection import train_test_split
from sklearn.metrics import roc_auc_score
clients = pd.read_csv("clients_churn.csv")
X = pd.get_dummies(clients[["anciennete", "appels_support", "montant", "contrat", "incidents_3m"]], dtype=float)
y = clients["churn"].values
X_train, X_test, y_train, y_test = train_test_split(X.values, y, test_size=0.3, random_state=42, stratify=y)
X_fit, X_valid, y_fit, y_valid = train_test_split(X_train, y_train, test_size=0.2, random_state=42, stratify=y_train)
# 3 étapes de décision ; arrêt quand l'AUC de validation ne progresse plus
modele = TabNetClassifier(n_d=8, n_a=8, n_steps=3, seed=42, verbose=0)
modele.fit(X_fit, y_fit, eval_set=[(X_valid, y_valid)], eval_metric=["auc"],
max_epochs=100, patience=15, batch_size=256, virtual_batch_size=128)
proba = modele.predict_proba(X_test)[:, 1]
print("AUC test :", round(roc_auc_score(y_test, proba), 3))
print(pd.Series(modele.feature_importances_, index=X.columns).sort_values(ascending=False).round(3))
# Masque d'attention : poids de chaque variable pour le premier client test
masque, _ = modele.explain(X_test[:1])
print(pd.Series(masque[0], index=X.columns).round(3))
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.
À 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.
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.
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 attentionLe 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 clientAttribue à 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 →Dataistudio forme les équipes au machine learning et à l'IA, sur des cas concrets.
Nous utilisons des cookies de mesure d'audience et de suivi publicitaire pour comprendre la fréquentation du site et l'efficacité de nos annonces. Rien n'est déposé sans votre accord. En savoir plus