メインコンテンツまでスキップ

TabFM: ゼロショット表形式基盤モデル

Open in Databricks

TabFM は、表形式データのための Google Research の基盤モデルです。ファインチューニング、ハイパーパラメータ探索、またはデータセット固有のトレーニングを必要とせず、トレーニング行がコンテキストとして渡されて1回のフォワードパスで予測が行われるインコンテキスト学習を使用します。数値列とカテゴリ列が混在するテーブルでの二項分類および多クラス分類(最大10クラス)と回帰をサポートしています。

このノートブックでは、乳がんデータセットでゼロショット分類を実行し、糖尿病データセットでゼロショット回帰を実行します。

注記

この例では、Databricks AI 環境バージョン 6 以上が必要です。

Serverless GPUコンピュートに接続する

Connect ドロップダウンをクリックし、 Serverless GPU を選択します。 Environment サイドパネルを開き、 Accelerator1xA10 に設定し、 AI v6 を選択します。

要件

  • 最初のラン時に Hugging Face Hub からモデルのウェイトをダウンロードするためのインターネットアクセス。
  • Databricks シークレットとして保存された Hugging Face 読み取りトークン。認証ステップの hf_secret_scope および hf_secret_key ウィジェットをシークレットのスコープとキーに設定します。
  • モデルの重みは、TabFM Non-Commercial License v1.0 の下でライセンスされています。
  • このノートブックには tabfm-1.0.0-pytorch のソースコードが含まれています。Copyright Google Research。Apache 2.0 ライセンスの下でライセンスされています。

TabFM は Databricks AI 環境バージョン 6 にプレインストールされているため、追加のインストールは必要ありません。

ライブラリのインポート

PyTorch、scikit-learn のデータセットローダーとメトリクス、および tabfm パッケージから TabFMClassifier / TabFMRegressor をインポートし、GPU の可用性を検証します。

Python
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 で認証する

hf_secret_scope および hf_secret_key ウィジェットを、Hugging Faceの読み取りトークンを格納するDatabricksのシークレットスコープとキーに設定し、Hubクライアントがdownloadの認証を行えるようにログインします。

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

ゼロショットの分類

乳がんデータセット(569 サンプル、30 個の数値特徴量)でゼロショット分類を実行します。mean radius から派生したカテゴリカルな radius_band 列が追加され、入力テーブルに数値型とカテゴリカル型が混在するようになります。TabFM はトレーニング行をコンテキストとして渡し、単一のフォワードパスでテストラベルを予測します。

分類データセットの読み込みと分割

乳がんデータセットをロードし、派生したカテゴリカル radius_band 特徴量を追加し、ターゲットで層化してトレーニング 80% / テスト 20% に分割します。

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

フィッティングと予測

分類モデルの重みをロードし、すべてのトレーニング行をインコンテキスト例として渡し、テストセットのクラスラベルと確率を予測します。精度と ROC-AUC を報告します。

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

ゼロショット回帰

糖尿病データセット(442サンプル、10個の数値特徴量)でゼロショット回帰を実行します。カテゴリ bmi_band 列が追加されます。TabFMは、各テストサンプルに対して連続的な疾病の進行スコアを予測します。

回帰データセットの読み込みと分割

糖尿病データセットを読み込み、派生したカテゴリカルな bmi_band 特徴量を追加し、80%をトレーニング用、20%をテスト用に分割します。

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

フィッティングと予測

回帰モデルの重みをロードし、すべてのトレーニング行をインコンテキスト例として渡し、テストセットの連続スコアを予測します。RMSE と R² を報告します。

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

次のステップ

このノートブックを別のデータセットに適応させるには、pandas DataFrame をロードし、ターゲットカラムを分離し、カテゴリカルカラムを文字列のままにし、トレーニングセットとテストセットに分割し、TabFMClassifier または TabFMRegressor にスワップします。TabFM はトレーニング行をインコンテキスト例として渡すため、メモリ使用量はトレーニングセットのサイズに応じてスケールします。したがって、大きなテーブルの場合は代表的なサンプルから始め、分類ターゲットを 10 クラス以下に維持してください。

サンプルノートブック

TabFM:ゼロショット表形式基盤モデル