Skip to main content

TF2 Keras implementation of TabNet

TabNet is a novel deep learning architecture for tabular data. TabNet performs reasoning in multiple decision steps and using sequential attention to select which features to use at which decision step. You can find more information about it in the original research paper.

Installation

$ pip install tabnet_keras

Usage

from tabnet_keras import TabNetRegressor, TabNetClassifier

tabnet_params = {
    "decision_dim": 16,
    "attention_dim": 16,
    "n_steps": 3,
    "n_shared_glus": 2,
    "n_dependent_glus": 2,
    "relaxation_factor": 1.3,
    "epsilon": 1e-15,
    "momentum": 0.98,
    "mask_type": "sparsemax", # can be 'sparsemax' or 'softmax'
    "lambda_sparse": 1e-3, 
    "virtual_batch_splits": 8 #number of splits for ghost batch normalization, ideally should evenly divide the batch_size
}

### Regression 
model = TabNetRegressor(n_regressors = 1, **tabnet_params)
model.compile(loss = 'mean_squared_error', optimizer = tf.keras.optimizers.Adam(0.01), 
             metrics = [tf.keras.metrics.RootMeanSquaredError()])
model.fit(X, y, epochs = 100, batch_size = 1024)

### Classification
model = TabNetClassifier(n_classes = 10, out_activation = None, **tabnet_params)
model.compile(loss = 'categorical_crossentropy', optimizer = tf.keras.optimizers.Adam(0.01))
model.fit(X, y, epochs = 100, batch_size = 1024)

Acknowledgment

Most of the code is taken with minor changes from this repository.

Metadata

Release files for tabnet-keras 1.2.0

For a detailed explanation of source distributions (sdists) and built distributions (wheels), please see the package formats documentation.

Source distribution (sdist)

Source distribution for tabnet-keras 1.2.0
File Size Uploaded
tabnet_keras-1.2.0.tar.gz 11.9 kB Details

Built distribution (wheel)

Table of built distributions (wheels) for tabnet-keras 1.2.0
File Interpreter ABI Platform
tabnet_keras-1.2.0-py3-none-any.whl Python 3 none any Details

Total release size: 28.8 kB

Release files / tabnet_keras-1.2.0.tar.gz

Download URL tabnet_keras-1.2.0.tar.gz
Size 11.9 kB
Tags Source
SHA-256 checksum
How to use checksums
1b975913cc85bd1f9d908d5cc3673202a25e73511c89ae80e697bc14565819c7
BLAKE2b-256 checksum
How to use checksums
95ded0287064c8f499788796efe13856389b75f74d9645456c9e36efacc19205
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/4.0.1 CPython/3.9.12

Release files / tabnet_keras-1.2.0-py3-none-any.whl

Download URL tabnet_keras-1.2.0-py3-none-any.whl
Size 16.9 kB
Tags Python 3
SHA-256 checksum
How to use checksums
784a53a8266ff06bf9a25c8097844a6faf98e91b963d4bc991ff5a993d22e2cc
BLAKE2b-256 checksum
How to use checksums
0009f821044b11e4550b79fab25428342f7087b1075d0d1ac99bb99ab2062dcd
Upload date
Uploaded using Trusted Publishing?
What is trusted publishing?
No
Uploaded via twine/4.0.1 CPython/3.9.12

Release history Release notifications | RSS feed

This release

1.2.0 This release

2 release files

1.1.0

2 release files

Anthropic, PBC Visionary sponsor Bloomberg Visionary sponsor Hudson River Trading Visionary sponsor Meta Visionary sponsor NVIDIA Visionary sponsor Microsoft Sustainability sponsor Depot Continuous Integration AWS Cloud computing and Security Sponsor Datadog Monitoring Fastly CDN Google Download Analytics Sentry Error logging StatusPage Status page