TabNetHyperparameters
VariantTabNet neural network. See `setup_TabNet`.
Properties
batch_sizeinteger | object≥ 1Batch size.
penaltynumber | object≥ 0Sparsity regularization penalty.
clip_valuenumber | object | nullGradient clip value.
lossstring | objectLoss function. auto = set from outcome type.
epochsinteger | object≥ 1Number of training epochs.
drop_lastboolean | objectDrop the last incomplete batch.
decision_widthinteger | object | null≥ 1Decision prediction layer width.
attention_widthinteger | object | null≥ 1Attention embedding width.
num_stepsinteger | object≥ 1Number of decision steps.
feature_reusagenumber | object≥ 0Feature reusage coefficient.
mask_typestring | objectMasking function.
one of
"sparsemax""entmax"virtual_batch_sizeinteger | object≥ 1Virtual batch size (ghost batch normalization).
valid_splitnumber | object≥ 0< 1Fraction of data used for (tabnet-internal) validation.
learn_ratenumber | object> 0Learning rate.
optimizerstringOptimizer name, resolved by the tabnet backend.
lr_schedulerstring | nullLearning-rate scheduler. NULL = none.
one of
"step""reduce_on_plateau"nulllr_decaynumber | object≥ 0≤ 1Learning rate decay.
step_sizeinteger | object≥ 1Learning rate scheduler step size.
checkpoint_epochsinteger | object≥ 1Checkpoint interval in epochs.
cat_emb_diminteger | object≥ 1Categorical embedding dimension.
num_independentinteger | object≥ 1Number of independent GLU layers at each encoder step.
num_sharedinteger | object≥ 1Number of shared GLU layers at each encoder step.
num_independent_decoderinteger | object≥ 1Number of independent GLU layers for pretraining.
num_shared_decoderinteger | object≥ 1Number of shared GLU layers for pretraining.
momentumnumber | object≥ 0Momentum for batch normalization.
pretraining_rationumber | object≥ 0≤ 1Ratio of features to mask during pretraining.
devicestringCompute device.
one of
"auto""cpu""cuda"importance_sample_sizeinteger | object | null≥ 1Sample size for importance calculation.
early_stopping_monitorstring | objectMetric monitored for early stopping.
one of
"auto""valid_loss""train_loss"early_stopping_tolerancenumber | object≥ 0Minimum relative improvement to reset the patience counter.
early_stopping_patienceinteger | object≥ 0Number of epochs without improvement before stopping.
num_workersinteger≥ 0Number of subprocesses for data loading.
skip_importancebooleanSkip importance calculation.
ifwboolean | objectInverse Frequency Weighting in classification.
Relationships
Used by