TabFM: sıfır örnekli tablosal temel model

TabFM , tablo veri için Google Research temel modelidir. Bağlam içi öğrenme kullanır; burada eğitim satırları bağlam olarak aktarılır ve tahminler tek bir ileri geçişle yapılır; ince ayarlama, hiperparametre araması veya veri setine özgü eğitim gerektirmez. İkili ve çok sınıflı sınıflandırmayı (10 sınıfa kadar) ve karışık sayısal ve kategorik sütunlu tablolarda regresyonu destekler.

Bu defter, meme kanseri veri setinde sıfır atış sınıflandırması ve diyabet veri setinde sıfır atış regresyonunu çalıştırıyor.

Note

Bu örnek için Databricks AI ortamı sürüm 6 veya üzeri gereklidir.

Sunucusuz GPU hesaplamasına bağlanma

Bağlan açılan menüsüne tıklayın ve Sunucusuz GPU'ya tıklayın. Çevre yan panelini açın, Accelerator'ı1xA10'a ayarlayın ve AI v6'yı seçin.

Requirements

  • İlk denemede Hugging Face Hub'dan model ağırlıklarını indirmek için internet erişimi.
  • Bir Hugging Face okuma tokenı, Databricks sırrı olarak saklanıyor. Kimlik doğrulama adımındaki hf_secret_key ve hf_secret_scope widget’larını, sırrınızın kapsamına ve anahtarına ayarlayın.
  • Model ağırlıkları TabFM Ticari Olmayan Lisans v1.0 kapsamında lisanslanmıştır.
  • Bu defter, Apache 2.0 lisansı altında lisanslanmış tabfm-1.0.0-pytorch kaynaklı kaynak kodu içerir; bu kod Google Research'e aittir.

TabFM, Databricks AI ortamı sürüm 6'da önceden yüklüdür, bu yüzden ek kurulum gerekmez.

Kitaplıkları içeri aktarma

PyTorch'u, scikit-learn veri seti yükleyicileri ve metriklerini ve tabfm paketinden TabFMClassifier / TabFMRegressor içe aktarın, ardından GPU kullanılabilirliğini doğrulayın.

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)}")

Hugging Face ile kimlik doğrulama

hf_secret_scope ve hf_secret_key widget'larını Hugging Face okuma tokeninizi saklayan Databricks gizli kapsamı ve anahtarına ayarlayın, sonra giriş yapın ki Hub istemcisi indirmeleri doğrulayabilir.

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)

Sıfır örnekli sınıflandırma

Meme kanseri veri setinde sıfır atış sınıflandırması yapın (569 örnek, 30 sayısal özellik). mean radius öğesinden türetilen kategorik bir radius_band sütunu eklenir, böylece giriş tablosu sayısal ve kategorik türlerin karışımını içerir. TabFM, eğitim satırlarını bağlam olarak geçer ve test etiketlerini tek bir ileri geçişte tahmin eder.

Sınıflandırma veri setini yükleyin ve bölünün

Meme kanseri veri setini yükleyin, türetilmiş kategorik bir radius_band özellik ekleyin ve hedefe göre tabakalandırma ile %80 eğitim / %20 test olarak bölün.

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],
})

Uyum ve tahmin et

Sınıflandırma modeli ağırlıklarını yükleyin, tüm eğitim satırlarını bağlam içi örnekler olarak geçirin ve test seti için sınıf etiketlerini ve olasılıklarını tahmin edin. Doğruluk ve ROC-AUC değerlerini raporla.

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))

Sıfır örnekli regresyon

Diyabet veri setinde sıfır atış regresyon yapın (442 örnek, 10 sayısal özellik). Kategorik bmi_band bir sütun eklenir. TabFM, her test örneği için sürekli hastalık ilerlemesi skoru öngörür.

Regresyon veri setini yükleyin ve bölünün

Diyabet veri setini yükleyin, türetilmiş kategorik bmi_band özelliği ekleyin ve onu %80 eğitim / %20 test olarak bölün.

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],
})

Uyum ve tahmin et

Regresyon modeli ağırlıklarını yükleyin, tüm eğitim satırlarını bağlam içi örnekler olarak iletin ve test kümesi için sürekli puanları tahmin edin. RMSE ve R²'yi bildirin.

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))

Sonraki Adımlar

Bu not defterini başka bir veri kümesine uyarlamak için bir pandas DataFrame yükleyin, hedef sütununu ayırın, kategorik sütunları dizge olarak bırakın, veriyi eğitim ve test kümelerine bölün ve yerine TabFMClassifier veya TabFMRegressor kullanın. TabFM, eğitim satırlarını bağlam içi örnekler olarak geçtiği için, bellek kullanımı eğitim kümesi boyutuyla ölçeklenir, bu yüzden büyük tablolar için temsilci bir örneklemle başlayın ve sınıflandırma hedeflerini 10 veya daha az sınıfta tutun.

Örnek defter

TabFM: sıfır örnekli tablosal temel model

Dizüstü bilgisayar al