Learning From Failure: Integrating Negative Examples when Fine-tuning
Fine-Tuning Strategies. By Simply Adding A Pre-
fix or suffix that tells the model whether to gen-
We Are The First To Demonstrate The Value Of Neg-
ative trajectories and their application in agent-
Introduction
An agent is a model that has the ability to interact
Plete Tasks In Narrow And Specialized Domains
1Code and data are available at: https://github.com/ Reason-Wang/NAT.
Gpt-4 (Achiam Et Al., 2023), Using Them As The
core of an agent system to process information and
Make Decisions (Gravitas, 2024; Yoheinakajima,
2024). This line of work has resulted in agent sys-
Tems That Are Able To Perform Much More Complex
and general tasks.
Source Llms Through Paid Apis, Raising Concerns
about cost, latency, and reproducibility. Addition-
Ally, Existing Llms Were Not Developed For Agent
use cases (e.g., generating actions or calling tools),
And Few-Shot Prompting Offers Only Limited Learn-
ing support (Chen et al., 2023).
Subsequent Work Has Explored Fine-Tuning Llms
as agents, typically in three stages: data collection, fine-tuning, and inference (Chen et al., 2023; Zeng et al., 2023; Yin et al., 2023; Qiao et al., 2024).
At The Data Collection Stage, A Powerful Llm Such
as GPT-4 is employed to interact with the environ-
Ment, And The Llm-Generated Outputs And Environ-
ment observations are collected as trajectories. In the fine-tuning stage, smaller models are fine-tuned using only successful trajectories. The fine-tuned
Passing The Performance Of The Original Llm (Yin
et al., 2023).
Tories That Do Not Successfully Complete The Task
(i.e. negative examples), using only successful tra- jectories (i.e. positive examples) in the fine-tuning
Stage (Zeng Et Al., 2023; Chen Et Al., 2023; Qiao
et al., 2024). However, in tasks demanding intricate
Discarded Negative Samples Can Exceed 60%, Lead-
ing to substantial data and computational resource wastage.
Through Fine-Tuning? And (2) How Can We Optimize
the use of negative examples to enhance agent per-
Examples, And Observe That Incorporating Negative
examples generally yields benefits. For the second
(Nat) Paradigm That Explicitly Tells The Model To
differentiate between correct and incorrect inter- actions by adding prefixes or suffixes. Our exper-
Iments Demonstrate That Nat Outperforms Tradi-
tional methods by solely using positive examples or
Naively Combining Positive And Negative Ones, En-
abling better fine-tuning for low-resource data. In addition, we conduct extensive experiments to ana- lyze the learned agents’ behavior after fine-tuning
We Are The First To Utilize Negative Examples In
agent training.
Information Akin To Positive Examples Across
various tasks and prompting strategies.
Trade-Off Between Useful Information And Er-
rors in negative examples.
Erful Llm As The Core Of The Agent System Without
fine-tuning (Sumers et al., 2023; Wu et al., 2023;
Ruan Et Al., 2023; Zhao Et Al., 2023). However,
LLMs are optimized to generate natural language.
To Make Them Capable Of Using Tools And Making
decisions, current work typically collects trajecto-
Ries Generated By Gpt-3.5/4, Then Uses These Tra-
jectories to fine-tune a smaller LLM (Zeng et al.,
2023; Chen Et Al., 2023; Qiao Et Al., 2024; Chen
et al., 2024; Zhang et al., 2024; Zhou et al., 2024).
By Gpt-4 On Agentbench (Liu Et Al., 2023B) Tasks,
and only keep samples that receive the best rewards. Chen et al. (2023) collect trajectories on question
Answering Tasks And Fine-Tune Models With Sam-
ples that correctly answer the question. Liu et al.
Work And A Complex Filtering Mechanism To Col-
lect fine-tuning datasets. Qiao et al. (2024) divide an agent into sub-agents with different functions. They then synthesize trajectories for the respective
Agents. However, They Still Only Use Samples With
the best rewards. A simple ablation study was done
By Zeng Et Al. (2023). However, None Of This Work
has investigated the effectiveness of negative sam-
Ples In Detail. Although Not Directly Comparable,
in Table 1, we provide the results of these methods and ours on several benchmarks for reference.
Learning From Negative Results Can Be Divided
into prompt-based and fine-tuning-based methods.
Ing Parameters. Madaan Et Al. (2023) Use Llms
to first generate an output and then refine the out- put iteratively, while Shinn et al. (2023) employ an evaluator to provide external feedback. Zhao et al.
(2023) let the agent compare successful and unsuc- cessful trajectories, and extract insights based on
Comparison. The Success Of These Methods Relies
on the quality of the evaluator used to analyze the trajectories. The performance of fine-tuning-based methods is less predictable since model weights are
Updated, And Less Work Has Been Done On This. Li
et al. (2023) propose a two-stage training paradigm
Fine-Tuned. Our Work Focuses On Fine-Tuning Llms
as agents and is much simpler and more effective.
And Then Describe Our Agent Framework. We Then
introduce the whole pipeline of NAT, including data
Collection, Negative-Aware Reformatting (Which Is
the core part of our method that differentiates it from others), fine-tuning, and inference. Figure 1 outlines previous methods and our NAT paradigm.
Motivation
Our idea is motivated by two considerations. First, humans learn from mistakes and failures. Failure is
Generate
Weng earns $12 an hour for babysitting. Yesterday,
Did She Earn? Please Generate A Solution That
**incorrectly** answers the question.
By This Person Born In January, 1903. Please
generate a solution that **correctly** answers the question.
Thought: I Need To Search For This Film And
its director, who was born in January 1903.
French Film Directed By Grigori Aleksandrov. The
film is also known as Sentimental Romance ...
(C)
Figure 1: An overview of previous methods and our NAT paradigm. (a) Data collection, where interactions between LLMs and environments (tools) are collected. (b) Data processing, where previous methods simply filter out negative examples, while we reformat trajectories by adding prompts to task queries based on whether they are positive or negative. (c) An example of reformated positive and negative trajectories. We omit the system prompts here.
29.6
Table 1: Comparison with methods from other papers. We report the best results reported in the corresponding papers.
often seen as a stepping stone to success, as it offers insights into what does not work and highlights ar- eas that require change or development. We believe
That Powerful Llms Can Also Learn These Valuable
lessons from unsuccessful trajectories.
And Llms May Learn Unwanted Errors From Nega-
tive trajectories if negative examples are incorpo- rated directly. Therefore, we add a prefix or suffix to the query to differentiate positive and negative examples, explicitly telling the model whether the following trajectories they learn are correct.
Tion Is Delineated As Follows. First, The Llm Is
provided with a system prompt that outlines (a) the specific task to be addressed (for instance, “solve
A Mathematical Problem”), (B) The Tools That Are
permissible for task execution, and (c) the expected
Action Space And Output Format (For Example, Fin-
ish[N] signifies that N is the final answer). We do
Not Provide System Prompts In Figure 1, For Sim-
plicity. Second, a query instance is introduced. We
Prompt The Model To Answer The Query In The React
(Yao et al., 2023) format, which consists of reason- ing texts (referred to as “thoughts”) and “actions”. Finally, during the interaction phase, the system ex-
Ecutes The Llm-Generated Actions Using The Prede-
fined tools, returns the resulting observations back
To The Llm, And Prompts For Subsequent Actions
until the finish action of the task is generated, or the interaction rounds exceed a pre-defined threshold. Naturally, the task-solving process yields interac- tion trajectories between the LLM and the environ- ment (i.e., tools in our framework).
Result. For The Two Question-Answering Tasks, We
design a search tool with the Serper 2 API. It takes a search query as input and returns the Google search results. We further re-rank the search results using
Mpnet (Song Et Al., 2020) And Dpr (Karpukhin
et al., 2020).
Aware Training Paradigm Here, Where Negative-
aware reformatting is the core part of the paradigm that enables better agent tuning.
Answers As Seed Data. We Then Use Gpt-3.53 To
generate trajectories three times, each with differ-
Ent Temperatures (0.2, 0.5, And 0.7). This Allows
us to gather a diverse range of positive and nega-
Tive Samples. By Comparing Predicted Answers And
ground truth answers, we can label each trajectory as positive or negative.
Positive Samples From Negative Samples During The
agent tuning process aids in teaching the model to
Discern Between Successful And Unsuccessful Out-
comes. We append a string suffix to tell the model
We Use The Reformat-
ted trajectories to fine-tune LLMs. The loss is com- puted only on the part of the text generated by the
(Zheng Et Al., 2023). During Inference, We Prompt
the fine-tuned agent using the prompt for positive examples only.
2Https://Serper.Dev/
3We use GPT-3.5-1106 version. Although GPT-4 has the potential to produce even higher quality data, we opted for GPT-3.5 due to cost considerations.
4The actual prompts that we use in our experiments are slightly more complex than those provided here.
Gsm8K, Asdiv (Miao Et Al., 2020), Svamp (Patel
et al., 2021), and MultiArith (Roy and Roth, 2015). For question answering, we collect trajectories and
2018) And Strategyqa (Geva Et Al., 2021), Respec-
tively. More details are provided in Appendix B.
We Primarily Compare Nat With Two
baselines. The Vanilla setting uses positive exam-
Ous Work (Zeng Et Al., 2023; Chen Et Al., 2023;
Qiao et al., 2024; Liu et al., 2024) has done. The
Second Setting Includes Negative Examples Without
adding any prefix or suffix, which we call Negative- Unaware Training (NUT).
2023). All The Models Are Fine-Tuned For 2 Epochs
with a batch size of 64. We use a cosine scheduler
With 3% Of Total Steps As The Warm-Up. The Maxi-
mum learning rate is set to 5 × 10−5. We train the
Model With 4×A100 Gpus With Deepspeed Zero
3 stage (Rajbhandari et al., 2019).
Ing (Nat) Not Only Outperform The Corresponding
model trained only on positive examples (Vanilla),
Lights The Value Of Nat In Data-Scarce Scenarios,
which is common for agent tuning. It is worth noting that previous work (Zeng et al.,
71.28
Table 2: Overall results for math tasks. Each block is a setting with a specific model and number of positive examples. The best results are bolded and second best results are underlined
Table 3: Results Of Llama-2-7B And 13B On Hot-
potQA. All results are reported as the mean score of 5 runs. For the 7B model, we report results using 1,500 negative samples; for the 13B models, we use 2k nega- tive samples. NAT-2 means we divide negative samples into 2 classes based on quality. We discuss this in Sec- tion 6.1.
Egyqa With 1000 Positive Samples And 500 Negative
samples. not contradict our findings: as we discuss in Sec-
Tions 5.1 And 5.2, Performance Is Determined By
both the quantity and quality of the negative data.
Results On Hotpotqa And Strategyqa. Here, Nat-2
is a variant of NAT, where we divide negative data into two classes and use different prompts for each,
Improves Performance By More Than 2% In Em And
6% in f1 score compared to no negative samples.
More Than 8% And About 3% Improvements Com-
pared to no negative samples and NUT, respectively. This suggests that our method is also effective for question-answering tasks.
Negative-Aware Training. Specifically, We Seek To
address the following questions: (1) Given a fixed
Model Gain From Negative Trajectories? (3) Are
all types of negative examples beneficial? and (4) What factors contribute to negative-aware training
Data For Our Experiments, The Analysis Is Done On
the math task.
Impact Of Training Sample Quantity
Our initial analysis focuses on the influence of neg- ative sample quantity. We maintain a constant num-
Justing The Negative Samples From 0 To 12K. The
results, depicted in Figure 2, illustrate the relation- ship between the quantity of negative data and av-
Samples Is About 11K In Both Cases. Due To Data
availability, we did not experiment with more neg- atives.
Number Of Positive Samples And Variable Number Of
negative samples.
10000 Negatives
Figure 3: Performance for a fixed number of negative samples (10k) and variable number of positive samples.
Based On Insights From Table 2 And Figure 2, We
hypothesize that the ideal ratio of negative samples is not fixed. Instead, it is influenced by two main factors: (1) the number of positive samples, as the improvements are larger for fewer positive samples; and (2) the intrinsic quality of the negative samples.
Regarding The First Point, We Hypothesize That
the marginal utility of negative samples diminishes
Test This Point, We Maintain A Constant Number Of
negative samples while varying the quantity of pos- itive samples from 0 to 5k. As depicted in Figure 3,
There Is A Diminishing Return On The Performance
added by negative samples as the count of positive samples rises. For the second point, we investigate the effects of negative data quality in Section 5.2.
We Sourced Negative Data From Various Models To
investigate the impact of negative data quality in
Nat. Specifically, We Consider The Data From Gpt-
3.5 as high-quality examples. In contrast, we gen-
Marith
Avg.
67.42
Table 5: LLaMA-2 7B model results trained with differ- ent quality negative data. We use 10k negative samples and experiment with 2k or 5k positive samples.
Nat W/ 2000 Positives
Figure 4: Perplexity for the model trained with 2k posi- tive samples and differing numbers of negative samples. The three dashed lines are perplexity computed on mod- els tuned with differing numbers of positive trajectories (without negatives).
resent low-quality data. For experiments, we paired 2k positive examples with 10k negative examples.
The Outcomes Presented In Table 5 Underscore The
critical role of data quality in NAT. In the 2k posi-
Sample Setting, The Improvements Are −6.20 And
+3.25, respectively.
The
learnable part in trajectories are thoughts and ac- tions, where thoughts are the reasoning on the cur- rent situation and planning for what to do next.
Actions are which tool to call and the input to the
Tool. We Analyze The Trajectories Of The Gsm8K
(Cobbe et al., 2021) test set generated by LLaMA-2-
And Nat Respectively. Table 6 Shows The Accuracy,
action error (the percent of incorrectly calling a
Ai Adaptive Learning
This project focuses on ai adaptive learning using modern AI and machine learning techniques. The content below is adapted from research literature and practical implementation notes.
We propose a novel high-performance and interpretable canon-
addition, unlike tree learning, DNNs enable gradient descent- ical deep tabular data learning architecture, TabNet. TabNet based end-to-end learning for tabular data which can have a uses sequential attention to choose which features to reason multitude of benefits: (i) efficiently encoding multiple data from at each decision step, enabling interpretability and more types like images along with tabular data; (ii) alleviating the efficient learning as the learning capacity is used for the most need for feature engineering, which is currently a key aspect
salient features. We demonstrate that TabNet outperforms in tree-based tabular data learning methods; (iii) learning other variants on a wide range of non-performance-saturated from streaming data and perhaps most importantly (iv) end- tabular datasets and yields interpretable feature attributions to-end models allow representation learning which enables plus insights into its global behavior. Finally, we demonstrate many valuable application scenarios including data-efficient
self-supervised learning for tabular data, significantly improv- domain adaptation (Goodfellow, Bengio, and Courville 2016), ing performance when unlabeled data is abundant. generative modeling (Radford, Metz, and Chintala 2015) and
Introduction We propose a new canonical DNN architecture for tabular
Deep neural networks (DNNs) have shown notable success data, TabNet. The main contributions are summarized as: efficiently encode the raw data into meaningful representa- enabling flexible integration into end-to-end learning. tions, fuel the rapid progress. One data type that has yet to 2. TabNet uses sequential attention to choose which fea- see such success with a canonical architecture is tabular data. tures to reason from at each decision step, enabling in-
Despite being the most common data type in real-world AI terpretability and better learning as the learning capacity (as it is comprised of any categorical and numerical features), is used for the most salient features (see Fig. 1). This under-explored, with variants of ensemble decision trees for each input, and unlike other instance-wise feature se- Why? First, because DT-based approaches have certain bene- and van der Schaar 2019), TabNet employs a single deep
fits: (i) they are representionally efficient for decision mani- learning architecture for feature selection and reasoning. folds with approximately hyperplane boundaries which are 3. Above design choices lead to two valuable properties: (i) common in tabular data; and (ii) they are highly interpretable TabNet outperforms or is on par with other tabular learn- in their basic form (e.g. by tracking decision nodes) and there ing models on various datasets for classification and re-
are popular post-hoc explainability methods for their ensem- gression problems from different domains; and (ii) TabNet ble form, e.g. (Lundberg, Erion, and Lee 2018) – this is an enables two kinds of interpretability: local interpretability important concern in many real-world applications; (iii) they that visualizes the importance of features and how they are fast to train. Second, because previously-proposed DNN are combined, and global interpretability which quantifies
architectures are not well-suited for tabular data: e.g. stacked the contribution of each feature to the trained model. convolutional layers or multi-layer perceptrons (MLPs) are 4. Finally, for the first time for tabular data, we show signif- vastly overparametrized – the lack of appropriate inductive icant performance improvements by using unsupervised bias often causes them to fail to find optimal solutions for tab- pre-training to predict masked features (see Fig. 2).
ular decision manifolds (Goodfellow, Bengio, and Courville
Why is deep learning worth exploring for tabular data?
One obvious motivation is expected performance improve- Feature selection: Feature selection broadly refers to judi- Copyright © 2021, Association for the Advancement of Artificial ciously picking a subset of features based on their useful-
Professional occupation related Investment related
Feedback from Feedback to
Feature selection Input processing Feature selection Input processing
previous step next step … …
Predicted output (whether the income level >$50k)
selection enables interpretability and better learning as the capacity is used for the most salient features. TabNet employs multiple decision blocks that focus on processing a subset of input features for reasoning. Two decision blocks shown as examples process features that are related to professional occupation and investments, respectively, in order to predict the income level.
Unsupervised pre-training Supervised fine-tuning
Age Cap. gain Education Occupation Gender Relationship Age Cap. gain Education Occupation Gender Relationship 5 2000 ? Exec-managerial F Wife 6 2000 Bachelors Exec-managerial M Husband 1 0 ? Farming-fishing M ? 2 0 High-school Farming-fishing M Unmarried
? 50 Doctorate Prof-specialty M Husband 4 50 Doctorate Prof-specialty M Husband 2 ? ? Handlers-cleaners F Wife 2 0 High-school Handlers-cleaners F Wife 5 3000 Bachelors ? ? Husband 5 3000 Bachelors Exec-managerial M Husband
3 0 Bachelors ? F ? 3 100 Bachelors Prof-specialty F Wife ? 0 High-school Armed-Forces ? Husband 2 0 High-school Armed-Forces M Husband
TabNet decoder Decision making
Age Cap. gain Education Occupation Gender Relationship Income > $50k
3 M False
level can be guessed from the occupation, or the gender can be guessed from the relationship. Unsupervised representation learning by masked self-supervised learning results in an improved encoder model for the supervised learning task.
ward selection and Lasso regularization (Guyon and Elisseeff performance with compact representations. 2003) attribute feature importance based on the entire training Tree-based learning: DTs are commonly-used for tabular data, and are referred as global methods. Instance-wise fea- data learning. Their prominent strength is efficient picking ture selection refers to picking features individually for each of global features with the most statistical information gain
to maximize the mutual information between the selected mance of standard DTs, one common approach is ensembling features and the response variable, and in (Yoon, Jordon, and to reduce variance. Among ensembling methods, random van der Schaar 2019) by using an actor-critic framework to forests (Ho 1998) use random subsets of data with randomly mimic a baseline while optimizing the selection. Unlike these, selected features to grow many trees. XGBoost (Chen and
sity in end-to-end learning – a single model jointly performs recent ensemble DT approaches that dominate most of the feature selection and output mapping, resulting in superior recent data science competitions. Our experimental results
!# + Softmax !" < % !" > % !# > & !# > &
ReLU ReLU &
$" !" − $" % −1 −$" !" + $" % −1 −1 $# !# − $# & % −1 −$# !# + $# & !"
FC FC
W: [$" , - $" , 0, 0] W: [0, 0, $# , - $# ] !" < % b: [-a $" , a $" , -1, -1] b: [-1, -1, -d $# , d $# ] !# < & !" > % !# < & [!" ] [!# ]
M: [1, 0] M: [0, 1]
(right). Relevant features are selected by using multiplicative sparse masks on inputs. The selected features are linearly transformed, and after a bias addition (to represent boundaries) ReLU performs region selection by zeroing the regions. Aggregation of multiple regions is based on addition. As C and C get larger, the decision boundary gets sharper.
for various datasets show that tree-based models can be out- constructs a sequential multi-step architecture, where each performed when the representation capacity is improved with step contributes to a portion of the decision based on the deep learning while retaining their feature selecting property. selected features; (iii) improves the learning capacity via non- Integration of DNNs into DTs: Representing DTs with linear processing of the selected features; and (iv) mimics
DNN building blocks as in (Humbird, Peterson, and McClar- ensembling via higher dimensions and more steps. ren 2018) yields redundancy in representation and ineffi- cient learning. Soft (neural) DTs (Wang, Aggarwal, and Liu Fig. 4 shows the TabNet architecture for encoding tabu- functions, instead of non-differentiable axis-aligned splits. mapping of categorical features with trainable embeddings.
However, losing automatic feature selection often degrades We do not consider any global feature normalization, but performance. In (Yang, Morillo, and Hospedales 2018), a soft merely apply batch normalization (BN). We pass the same D- binning function is proposed to simulate DTs in DNNs, by dimensional features f ∈ <B×D to each decision step, where 2019) proposes a DNN architecture by explicitly leveraging multi-step processing with Nsteps decision steps. The ith
expressive feature combinations, however, learning is based step inputs the processed information from the (i − 1)th step on transferring knowledge from gradient-boosted DT. (Tanno to decide which features to use and outputs the processed ing from primitive blocks while representation learning into sion. The idea of top-down attention in the sequential form edges, routing functions and leaf nodes. TabNet differs from is inspired by its applications in processing visual and text
these as it embeds soft feature selection with controllable data (Hudson and Manning 2018) and reinforcement learn- Self-supervised learning: Unsupervised representation relevant information in high dimensional input. learning improves supervised learning especially in small Feature selection: We employ a learnable mask M[i] ∈ has shown significant advances – driven by the judicious capacity of a decision step is not wasted on irrelevant
choice of the unsupervised learning objective (masked input ones, and thus the model becomes more parameter effi- prediction) and attention-based deep learning. cient. The masking is multiplicative, M[i] · f . We use an attentive transformer (see Fig. 4) to obtain the masks us- TabNet for Tabular Learning ing the processed features from the preceding step, a[i − 1]:
M[i] = sparsemax(P[i − 1] · hi (a[i − 1])). Sparsemax nor-
DTs are successful for learning from real-world tabular malization (Martins and Astudillo 2016) encourages sparsity datasets. With a specific design, conventional DNN building by mapping the Euclidean projection onto the probabilistic blocks can be used to implement DT-like output manifold, simplex, which is observed to be superior in performance and e.g. see Fig. 3). In such a design, individual feature selec- aligned with the goal of sparse feature selection for explain-
tion is key to obtain decision boundaries in hyperplane form, PD which can be generalized to a linear combination of features ability. Note that j=1 M[i]b,j = 1. hi is a trainable func- where coefficients determine the proportion of each feature. tion, shown in Fig. 4 using a FC layer, followed by BN. P[i] TabNet is based on such functionality and it outperforms DTs is the prior scale term, denoting how much a particular feature
Qi while reaping their benefits by careful design which: (i) uses has been used previously: P[i] = j=1 (γ − M[j]), where γ sparse instance-wise feature selection learned from data; (ii) is a relaxation parameter – when γ = 1, a feature is enforced
+ Softmax
Feature Feature …
transformer transformer
x Nsteps Features
+ Softmax
Feature …
transformer transformer Feature Feature Feature Feature transformer
Encoded representation
transformer transformer Attentive transformer … Mask transformer …
Step 2 Decision step dependent
transformer transformer
BN Feature Feature
FC BN transformer transformer
+ 0.5 0.5 0.5 Agg. Agg. Features Features FC FC + +
Reconstructed + … Feature attributes + … features
(a) TabNet encoder architecture (b) TabNet decoder architecture Feature transformer Feature Attentive transformer Shared across decision steps Decision step dependent transformer GLU
Decision step dependent Prior scales
+ 0.5 0.5 0.5
0.5 0.5 0.5
+ Attentive transformer (c) (d)
Prior scales
divides the processed representation to be used by the attentive transformer of the subsequent step as well as for the overall Attentive BN FC
output. For each step, the feature selection mask provides interpretable information about the model’s functionality, and the +
masks can be aggregated to obtain global feature transformer important attribution. (b) TabNet decoder, composed of a feature transformer block at each step. (c) A feature transformer block example – 4-layer network is shown, where 2 are shared across all decision
Prior scales
steps and 2 are decision step-dependent. Each layer is composed of a fully-connected (FC) layer, BN and GLU nonlinearity. (d) +
An attentive transformer block example – a single layer mapping is modulated with a prior scale information which aggregates Sparsemax
how much each feature has been used before the current decision step. sparsemax (Martins and Astudillo 2016) is used for BN FC
normalization of the coefficients, resulting in sparse selection of the salient features. +
to be used only at one decision step and as γ increases, more propose the aggregate.feature importance mask, Magg−b,j = flexibility is provided to use a feature at multiple decision PNsteps ηb [i]Mb,j [i]
PD PNsteps
ηb [i]Mb,j [i].2 i=1 i=1 steps. P is initialized as all ones, 1B×D , without any prior j=1
on the masked features. If some features are unused (as in self- Tabular self-supervised learning: We propose a decoder supervised learning), corresponding P entries are made 0 architecture to reconstruct tabular features from the Tab- to help model’s learning. To further control the sparsity of the Net encoded representations. The decoder is composed of selected features, we propose sparsity regularization in the feature transformers, followed by FC layers at each deci-
form of entropy (Grandvalet and Bengio 2004), Lsparse = sion step. The outputs are summed to obtain the recon-
PNsteps PB PD −Mb,j [i] log(Mb,j [i]+)
i=1 b=1 j=1 Nsteps ·B , where is a structed features. We propose the task of prediction of miss- small number for numerical stability. We add the sparsity reg- ing feature columns from the others. Consider a binary mask ularization to the overall loss, with a coefficient λsparse . Spar- S ∈ {0, 1}B×D . The TabNet encoder inputs (1 − S) · f̂ sity provides a favorable inductive bias for datasets where and the decoder outputs the reconstructed features, S · f̂ . We
most features are redundant. initialize P = (1 − S) in encoder so that the model em- Feature processing: We process the filtered features using phasizes merely on the known features, and the decoder’s last a feature transformer (see Fig. 4) and then split for the FC layer is multiplied with S to output the unknown features. decision step output and information for the subsequent We consider the reconstruction loss in self-supervised phase:
step, [d[i], a[i]] = fi (M[i] · f ), where d[i] ∈ <B×Nd and 2
PB PD (f̂b,j −fb,j )·Sb,j
a[i] ∈ <B×Na . For parameter-efficient and robust learning b=1 j=1
√ PB PB 2
. Normalization b=1 (fb,j −1/B b=1 fb,j ) with high capacity, a feature transformer should comprise layers that are shared across all decision steps (as the same with the population standard deviation of the ground truth features are input across different decision steps), as well as is beneficial, as the features may have different ranges. We decision step-dependent layers. Fig. 4 shows the implementa- sample Sb,j independently from a Bernoulli distribution with
tion as concatenation of two shared layers and two decision parameter ps , at each iteration. step-dependent layers. Each FC layer is followed by BN and eventually connected to a normalized residual √ connection We study TabNet in wide range of problems, that contain with normalization. Normalization with 0.5 helps to sta- regression or classification tasks, particularly with published bilize learning by ensuring that the variance throughout the benchmarks. For all datasets, categorical inputs are mapped
For faster training, we use large batch sizes with BN. Thus, bedding and numerical columns are input without and pre- except the one applied to the input features, we use ghost BN processing.4 We use standard classification (softmax cross (Hoffer, Hubara, and Soudry 2017) form, using a virtual batch entropy) and regression (mean squared error) loss functions size BV and momentum mB . For the input features, we ob- and we train until convergence. Hyperparameters of the Tab-
serve the benefit of low-variance averaging and hence avoid Net models are optimized on a validation set and listed in ghost BN. Finally, inspired by decision-tree like aggregation Appendix. TabNet performance is not very sensitive to most as in Fig. 3, we construct the overall decision embedding hyperparameters as shown with ablation studies in Appendix. as dout = i=1 PNsteps ReLU(d[i]). We apply a linear mapping In Appendix, we also present ablation studies on various de-
Wfinal dout to get the output mapping.1 sign and guidelines on selection of the key hyperparameters. Interpretability: TabNet’s feature selection masks can shed For all experiments we cite, we use the same training, val- light on the selected features at each step. If Mb,j [i] = 0, idation and testing data split with the original work. Adam optimization algorithm (Kingma and Ba 2014) and Glorot then j th feature of the bth sample should have no contribution uniform initialization are used for training of all models.5
to the decision. If fi were a linear function, the coefficient
Mb,j [i] would correspond to the feature importance of fb,j . Instance-wise feature selection
Although each decision step employs non-linear processing, their outputs are combined later in a linear way. We aim Selection of the salient features is crucial for high perfor- to quantify an aggregate feature importance in addition to mance, especially for small datasets. We consider 6 tabular requires a coefficient that can weigh the relative importance samples). The datasets are constructed in such a way that of each step in the decision. We simply propose ηb [i] = only a subset of the features determine the output. For Syn1-
PNd Syn3, salient features are same for all instances (e.g., the
c=1 ReLU(db,c [i]) to denote the aggregate decision con- tribution at ith decision step for the bth sample. Intuitively, if 2
Normalization is used to ensure D
P j=1 Magg−b,j = 1. db,c [i] < 0, then all features at ith decision step should have 3
0 contribution to the overall decision. As its value increases, prove the performance, but interpretation of individual dimensions
it plays a higher role in the overall linear combination. Scal- may become challenging. ing the decision mask at each decision step with ηb [i], we Specially-designed feature engineering, e.g. logarithmic trans- formation of variables highly-skewed distributions, may further
For discrete outputs, we additionally employ softmax during
training (and argmax during inference). An open-source implementation will be released.
Global: using only globally-salient features, Tree Ensembles (Geurts, Ernst, and Wehenkel 2006), Lasso-regularized model, L2X
Syn Syn Syn Syn Syn Syn
No selection .5 ± .0 .7 ± .0 .8 ± .0 .5 ± .0 .6 ± .0 .6 ± .0 Tree .5 ± .1 .8 ± .0 .8 ± .0 .6 ± .0 .7 ± .0 .7 ± .0 Lasso-regularized .4 ± .0 .5 ± .0 .8 ± .0 .5 ± .0 .6 ± .0 .7 ± .0
INVASE .6 ± .0 .8 ± .0 .9 ± .0 .7 ± .0 .7 ± .0 .8 ± .0
Global .6 ± .0 .8 ± .0 .9 ± .0 .7 ± .0 .7 ± .0 .8 ± .0 TabNet .6 ± .0 .8 ± .0 .8 ± .0 .7 ± .0 .7 ± .0 .8 ± .0
output of Syn depends on features X -X ), and global fea- Table 3: Performance for Poker Hand induction dataset. ture selection, as if the salient features were known, would give high performance. For Syn4-Syn6, salient features are Model Test accuracy (%) instance dependent (e.g., for Syn4, the output depends on ei- DT 50.0 ther X -X or X -X depending on the value of X ), which MLP 50.0
makes global feature selection suboptimal. Table 1 shows that Deep neural DT 65.1
TabNet outperforms others (Tree Ensembles (Geurts, Ernst, XGBoost 71.1
and Wehenkel 2006), LASSO regularization, L2X (Chen LightGBM 70.0 van der Schaar 2019). For Syn1-Syn3, TabNet performance TabNet 99.2 is close to global feature selection - it can figure out what Rule-based 100.0 features are globally important. For Syn4-Syn6, eliminating instance-wise redundant features, TabNet improves global feature selection. All other methods utilize a predictive model Poker Hand (Dua and Graff 2017): The task is classifica-
with 43k parameters, and the total number of parameters is tion of the poker hand from the raw suit and rank attributes of 101k for INVASE due to the two other models in the actor- the cards. The input-output relationship is deterministic and critic framework. TabNet is a single architecture, and its size hand-crafted rules can get 100% accuracy. Yet, conventional is 26k for Syn1-Syn and 31k for Syn4-Syn6. The compact DNNs, DTs, and even their hybrid variant of deep neural DTs
representation is one of TabNet’s valuable properties. (Yang, Morillo, and Hospedales 2018) severely suffer from the imbalanced data and cannot learn the required sorting and Performance on real-world datasets ranking operations (Yang, Morillo, and Hospedales 2018).
Tuned XGBoost, CatBoost, and LightGBM show very slight
as it can perform highly-nonlinear processing with its depth, Model Test accuracy (%) without overfitting thanks to instance-wise feature selection.
CatBoost 85.1 Table 4: Performance on Sarcos dataset. Three TabNet mod-
AutoML Tables 94.9 els of different sizes are considered.
Forest Cover Type (Dua and Graff 2017): The task is clas- MLP 2.1 0.14M
sification of forest cover type from cartographic variables. Adaptive neural tree 1.2 0.60M approaches that are known to achieve solid performance (AutoML 2019), an automated search framework based on TabNet-M 0.2 0.59M ensemble of models including DNN, gradient boosted DT, TabNet-L 0.1 1.75M with very thorough hyperparameter search. A single TabNet without fine-grained hyperparameter search outperforms it. Sarcos (Vijayakumar and Schaal 2000): The task is re-
gressing inverse dynamics of an anthropomorphic robot arm.
very small model is possible with a random forest. In the very and TabNet merely focuses on the relevant ones. For Syn4, small model size regime, TabNet’s performance is on par the output depends on either X -X or X -X depending parameters. When the model size is not constrained, TabNet feature selection – it allocates a mask to focus on the indi- achieves almost an order of magnitude lower test MSE. cator X , and assigns almost all-zero weights to irrelevant
features (the ones other than two feature groups). models are denoted with -S and -M. Real-world datasets: We first consider the simple task of mushroom edibility prediction (Dua and Graff 2017). Tab- Model Test acc. (%) Model size Net achieves 100% test accuracy on this dataset. It is indeed Sparse evolutionary MLP 78.4 81K known (Dua and Graff 2017) that “Odor” is the most discrim-
What is this project about?
This project covers practical implementation and research aspects of the topic using AI/ML techniques.