Skip to main content

TabFM: Tabular Foundation Models

TabFM (Tabular Foundation Model) is a scikit-learn compatible tabular foundation model. It allows you to perform zero-shot classification and regression on tabular datasets with mixed column types out-of-the-box.

At inference time, TabFM does not require training parameters on your dataset; instead, it leverages in-context learning by reading your training data as "context" to make instant predictions on new test samples.

This is not an officially supported Google product.


Installation

To install TabFM, clone the repository and install it locally with the backend of your choice:

JAX (CPU):

git clone https://github.com/google-research/tabfm.git
cd tabfm
pip install -e .[jax]

JAX (GPU):

git clone https://github.com/google-research/tabfm.git
cd tabfm
pip install -e .[jax,cuda]

PyTorch (CPU/GPU):

git clone https://github.com/google-research/tabfm.git
cd tabfm
pip install -e .[pytorch]

Note: For PyTorch with GPU support, ensure you have the appropriate PyTorch version installed for your CUDA version before installing TabFM.

Requirements

For a complete list of pinned dependencies and versions, please see requirements.txt. The core requirements depend on the backend you choose:

  • Python >= 3.11
  • Hugging Face Hub (for downloading pre-trained weights)
  • JAX Backend:
    • JAX (specifically jax==0.10.1)
    • Flax (specifically flax==0.12.7, using the modern flax.nnx API)
  • PyTorch Backend:
    • PyTorch (specifically torch==2.12.1+cpu or a GPU version)

Quick Start (TabFM v1.0.0)

We provide pre-trained weights for the TabFM v1.0.0 release. The library handles downloading and loading these weights automatically. You can choose to load the model using either the JAX or PyTorch backend.

1. Classification Example

import numpy as np
import pandas as pd
from tabfm import TabFMClassifier

# Choose your backend:

# OPTION A: JAX Backend
from tabfm import tabfm_v1_0_0_jax as tabfm_v1_0_0
model = tabfm_v1_0_0.load()

# OPTION B: PyTorch Backend
# from tabfm import tabfm_v1_0_0_pytorch as tabfm_v1_0_0
# model = tabfm_v1_0_0.load()

# Initialize scikit-learn compatible classifier (works with either backend model)
clf = TabFMClassifier(model=model)

# Prepare your dataset (supports mixed numerical and categorical features)
X_train = pd.DataFrame({
    "age": [25.0, 45.0, 35.0, 50.0],
    "job": ["engineer", "manager", "engineer", "manager"],
    "income": [80000, 120000, 90000, 130000]
})
y_train = np.array(["low_risk", "high_risk", "low_risk", "high_risk"])

X_test = pd.DataFrame({
    "age": [30.0, 48.0],
    "job": ["engineer", "manager"],
    "income": [85000, 125000]
})

# Fit classifier (prepares ordinal encoders and numerical scalers)
clf.fit(X_train, y_train)

# Predict classes and probabilities
predictions = clf.predict(X_test)
probabilities = clf.predict_proba(X_test)

print("Predictions:", predictions)
print("Class Probabilities:\n", probabilities)

2. Regression Example

import numpy as np
import pandas as pd
from tabfm import TabFMRegressor

# Choose your backend:

# OPTION A: JAX Backend
from tabfm import tabfm_v1_0_0_jax as tabfm_v1_0_0
model = tabfm_v1_0_0.load(model_type="regression")

# OPTION B: PyTorch Backend
# from tabfm import tabfm_v1_0_0_pytorch as tabfm_v1_0_0
# model = tabfm_v1_0_0.load(model_type="regression")

# Initialize scikit-learn compatible regressor (works with either backend model)
reg = TabFMRegressor(model=model)

# Prepare your dataset
X_train = pd.DataFrame({
    "sqft": [1200, 2500, 1500, 3000],
    "neighborhood": ["A", "B", "A", "C"]
})
y_train = np.array([250000, 550000, 310000, 620000])

X_test = pd.DataFrame({
    "sqft": [1800, 2800],
    "neighborhood": ["A", "B"]
})

# Fit and Predict
reg.fit(X_train, y_train)
predictions = reg.predict(X_test)

print("Predicted Prices:", predictions)

Examples Directory

You can find runnable scripts for both classification and regression under the examples/ folder:

To run them, simply execute:

python examples/classification_example.py

(You can edit these files to switch between JAX and PyTorch backends as shown in the comments inside them).


Evaluation Results

Our model evaluation results can be found in results/.


Running Tests

You can run the unit tests directly using Python's unittest module:

# Run all tests (requires both JAX and PyTorch installed)
PYTHONPATH=. python3 -m unittest discover -s tabfm/src/ -p "*_test.py"

# Or run specific test files:
PYTHONPATH=. python3 -m unittest tabfm/src/pytorch/model_test.py
PYTHONPATH=. python3 -m unittest tabfm/src/classifier_and_regressor_pytorch_test.py

Alternatively, if you have Bazel installed, you can run tests with:

bazel test //...

Download files

Download the file for your platform. If you're not sure which to choose, learn more about installing packages.

Source Distribution

tabfm-1.0.1.tar.gz (74.9 kB view details)

Uploaded Source

Built Distribution

If you're not sure about the file name format, learn more about wheel file names.

tabfm-1.0.1-py3-none-any.whl (81.0 kB view details)

Uploaded Python 3

File details

Details for the file tabfm-1.0.1.tar.gz.

File metadata

  • Download URL: tabfm-1.0.1.tar.gz
  • Upload date:
  • Size: 74.9 kB
  • Tags: Source
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.12.13

File hashes

Hashes for tabfm-1.0.1.tar.gz
Algorithm Hash digest
SHA256 77f0d871a6772560f1b0744c1822ac89f9c1eabe80768ac355df1aa48e116fbe
MD5 4f82f3b51ab20b12cfd61241a7da5270
BLAKE2b-256 d99ac52b3d7e23ef8a79883cc1bedf749ba519fb2686c5a3d0c43ff704e0aad8

See more details on using hashes here.

File details

Details for the file tabfm-1.0.1-py3-none-any.whl.

File metadata

  • Download URL: tabfm-1.0.1-py3-none-any.whl
  • Upload date:
  • Size: 81.0 kB
  • Tags: Python 3
  • Uploaded using Trusted Publishing? No
  • Uploaded via: twine/6.2.0 CPython/3.12.13

File hashes

Hashes for tabfm-1.0.1-py3-none-any.whl
Algorithm Hash digest
SHA256 79f5873c0e44bc853d00e8e38e33b8263fa63c5408ad63833ebd60bab91df99e
MD5 ad04b78e7970e7108986821d2db6ba8f
BLAKE2b-256 9e8350c0e4c8fb2a333075ae1419658a91a68453772a3f37c35fa0e8a07a5f81

See more details on using hashes here.

Supported by

AWS Cloud computing and Security Sponsor Datadog Monitoring Depot Continuous Integration Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page