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.
- Code: https://github.com/mrbakhtyari/TabK
- Paper: link will be added once the paper is published
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