TabK

Pretrained weights for TabK: Amortized Bayesian Estimation of the Number of Clusters in Tabular Data (NeurIPS 2026).

TabK is a permutation-invariant Transformer that estimates the number of clusters k of a tabular dataset in a single forward pass, without running a clustering algorithm.

Usage

git clone --depth 1 https://github.com/mrbakhtyari/TabK.git
cd TabK
uv sync
from sklearn.datasets import load_iris
from tabk import TabK

X, _ = load_iris(return_X_y=True)

model = TabK.from_pretrained()  # downloads this repository
print(model.predict(X))         # 3

Or from the command line:

uv run tabk predict data.csv

Files

File Description
fold_1.safetensors โ€ฆ fold_5.safetensors Weights of the 5-fold ensemble; predictions are averaged across folds
config.json Model and head configuration shared by all folds

Model details

  • Output: k โˆˆ {2, โ€ฆ, 15}, via a DLDL (label distribution) head
  • Training data: 40,000 synthetic datasets from a diverse generative prior, with 100โ€“2,500 rows and 2โ€“200 features
  • Parameters: 4.4M per fold (d_model 256, 8 heads, 4 layers)
  • Preprocessing: features are standardized automatically by TabK.predict
  • Training compute: ~37.6 hours on a single NVIDIA A100-SXM4-40GB

Limitations

  • Predictions are restricted to 2โ€“15 clusters.
  • Tables with more than 2,500 rows are uniformly subsampled to 2,500 rows; more than 200 features is outside the training range.
  • Inputs are expected to be numeric, with no missing values.

License

MIT

Downloads last month
45
Inference Providers NEW
This model isn't deployed by any Inference Provider. ๐Ÿ™‹ Ask for provider support