TabFM: zero-shot tabular foundation model
TabFM est un modèle de fondation de Google Research pour les données tabulaires. Il utilise l'apprentissage en contexte, où les lignes d'entraînement sont transmises comme contexte et les prédictions sont faites en une seule passe avant, sans affinement, recherche d'hyperparamètres ou entraînement spécifique au dataset requis. Il prend en charge la classification binaire et multiclasse (jusqu'à 10 classes) ainsi que la régression sur des tableaux comportant des colonnes numériques et catégorielles mixtes.
Ce notebook exécute une classification zero-shot sur le dataset du cancer du sein et une régression zero-shot sur le dataset du diabète.
This example requires the Databricks AI environment version 6 or above.
Connect to serverless GPU compute
Cliquez sur le menu déroulant Connect et sélectionnez GPU serverless . Ouvrez le panneau latéral Environment , définissez Accelerator sur 1xA10 et sélectionnez AI v6 .
Exigences
- Internet access to download the model weights from the Hugging Face Hub on the first run.
- A Hugging Face read token stored as a Databricks secret. Set the
hf_secret_scopeandhf_secret_keywidgets in the authentication step to your secret's scope and key. - Les poids du modèle sont sous licence TabFM Non-Commercial License v1.0.
- Ce notebook inclut le code source de tabfm-1.0.0-pytorch, Copyright Google Research, sous licence Apache 2.0.
TabFM est préinstallé dans l'environnement Databricks AI version 6, de sorte qu'aucune installation supplémentaire n'est requise.
Import libraries
Import PyTorch, the scikit-learn dataset loaders and metrics, and TabFMClassifier / TabFMRegressor from the tabfm package, then verify GPU availability.
import numpy as np
import pandas as pd
import torch
from sklearn.datasets import load_breast_cancer, load_diabetes
from sklearn.metrics import accuracy_score, roc_auc_score, mean_squared_error, r2_score
from sklearn.model_selection import train_test_split
from tabfm import TabFMClassifier, TabFMRegressor, tabfm_v1_0_0_pytorch as tabfm_v1_0_0
print(f"Torch version: {torch.__version__}")
print(f"CUDA available: {torch.cuda.is_available()}")
if torch.cuda.is_available():
print(f"GPU: {torch.cuda.get_device_name(0)}")
S'authentifier auprès de Hugging Face
Définissez les widgets hf_secret_scope et hf_secret_key sur le Secret Scope et la clé Databricks qui stockent votre jeton de lecture Hugging Face, puis connectez-vous pour que le client Hub puisse authentifier les téléchargements.
from huggingface_hub import login
# Set these widgets to the Databricks secret scope and key that hold your Hugging Face read token.
dbutils.widgets.text("hf_secret_scope", "", "Hugging Face secret scope")
dbutils.widgets.text("hf_secret_key", "hf_token", "Hugging Face secret key")
hf_token = dbutils.secrets.get(
scope=dbutils.widgets.get("hf_secret_scope"),
key=dbutils.widgets.get("hf_secret_key"),
)
login(token=hf_token)
Classification zero-shot
Exécutez la classification zero-shot sur le dataset du cancer du sein (569 échantillons, 30 caractéristiques numériques). Une colonne radius_band catégorielle dérivée de mean radius est ajoutée afin que la table d'entrée associe des types numériques et catégoriels. TabFM transmet les lignes d'entraînement en tant que contexte et prédit les étiquettes de test en une seule passe avant.
Charger et diviser le dataset de classification
Load the breast cancer dataset, add a derived categorical radius_band feature, and split it 80% train / 20% test with stratification on the target.
breast = load_breast_cancer(as_frame=True)
clf_df = breast.frame.copy()
clf_df["radius_band"] = pd.qcut(
clf_df["mean radius"],
q=4,
labels=["small", "medium", "large", "xlarge"],
).astype(str)
X_clf = clf_df.drop(columns=["target"])
y_clf = clf_df["target"]
X_train_clf, X_test_clf, y_train_clf, y_test_clf = train_test_split(
X_clf,
y_clf,
test_size=0.2,
random_state=42,
stratify=y_clf,
)
display(X_train_clf.head(5))
print({
"train_rows": len(X_train_clf),
"test_rows": len(X_test_clf),
"feature_count": X_train_clf.shape[1],
})
Ajustement et prédiction
Load the classification model weights, pass all training rows as in-context examples, and predict class labels and probabilities for the test set. Report accuracy and ROC-AUC.
tabfm_clf_model = tabfm_v1_0_0.load(model_type="classification")
tabfm_clf = TabFMClassifier(model=tabfm_clf_model)
tabfm_clf.fit(X_train_clf, y_train_clf)
clf_pred_proba = np.asarray(tabfm_clf.predict_proba(X_test_clf))
clf_pred = np.asarray(tabfm_clf.predict(X_test_clf)).reshape(-1)
clf_results = pd.DataFrame({
"actual": y_test_clf.reset_index(drop=True),
"predicted": clf_pred.astype(int),
"positive_class_probability": clf_pred_proba[:, 1],
})
accuracy = accuracy_score(y_test_clf, clf_pred)
roc_auc = roc_auc_score(y_test_clf, clf_pred_proba[:, 1])
print({
"accuracy": round(float(accuracy), 4),
"roc_auc": round(float(roc_auc), 4),
})
display(clf_results.head(10))
Zero-shot regression
Exécutez la régression zero-shot sur le dataset de diabète (442 échantillons, 10 caractéristiques numériques). Une colonne catégorielle bmi_band est ajoutée. TabFM prédit un score de progression continue de la maladie pour chaque échantillon de test.
Charger et fractionner le dataset de régression
Charger le dataset sur le diabète, ajouter une feature bmi_band catégorielle dérivée et effectuer une division de 80 % pour l'entraînement / 20 % pour le test.
diabetes = load_diabetes(as_frame=True)
reg_df = diabetes.frame.copy()
reg_df["bmi_band"] = pd.qcut(
reg_df["bmi"],
q=4,
labels=["low", "mid_low", "mid_high", "high"],
).astype(str)
X_reg = reg_df.drop(columns=["target"])
y_reg = reg_df["target"]
X_train_reg, X_test_reg, y_train_reg, y_test_reg = train_test_split(
X_reg,
y_reg,
test_size=0.2,
random_state=42,
)
display(X_train_reg.head(5))
print({
"train_rows": len(X_train_reg),
"test_rows": len(X_test_reg),
"feature_count": X_train_reg.shape[1],
})
Ajustement et prédiction
Load the regression model weights, pass all training rows as in-context examples, and predict continuous scores for the test set. Report RMSE and R².
tabfm_reg_model = tabfm_v1_0_0.load(model_type="regression")
tabfm_reg = TabFMRegressor(model=tabfm_reg_model)
tabfm_reg.fit(X_train_reg, y_train_reg)
reg_pred = np.asarray(tabfm_reg.predict(X_test_reg)).reshape(-1)
rmse = np.sqrt(mean_squared_error(y_test_reg, reg_pred))
r2 = r2_score(y_test_reg, reg_pred)
reg_results = pd.DataFrame({
"actual": y_test_reg.reset_index(drop=True),
"predicted": reg_pred,
})
reg_results["absolute_error"] = (reg_results["actual"] - reg_results["predicted"]).abs()
print({
"rmse": round(float(rmse), 4),
"r2": round(float(r2), 4),
})
display(reg_results.head(10))
Étapes suivantes
To adapt this notebook to another dataset, load a pandas DataFrame, separate the target column, leave categorical columns as strings, split into train and test sets, and swap in TabFMClassifier or TabFMRegressor. Because TabFM passes training rows as in-context examples, memory usage scales with training-set size, so start with a representative sample for large tables, and keep classification targets at 10 or fewer classes.
- TabFM model card
- Meilleures pratiques pour le compute serverless GPU
- Résoudre les problèmes sur le compute GPU serverless
- Machine Learning classique sur GPU serverless
- AI Runtime example notebooks