Enquire Now
Medical Computer Vision · Clinical Diagnostics · PyTorch / TensorFlow · GPU Optimized · 2026

Generative Ai Image Generation Capstone Project

Tensor Pipeline · Custom Loss Formulations · Model Quantization · Accelerated Inference — A rigorous deep learning engineering project focused on automated pathological lesion segmentation and radiological disease classification. Architected for thesis defense viva presentations, IEEE reproduction, and high-throughput production deployment.

PyTorch
Core Framework
AMP FP16
Mixed Precision
TensorRT
Quantized Serving

Qingyao Ai1†, Jingtao Zhan1†, Yiqun Liu1

Beijing, China. †These authors contributed equally to this work.

Abstract

The chapter discusses the foundational impact of modern generative AI models on information access (IA) systems. In contrast to traditional AI, the large-scale training and superior data modeling of generative AI models enable them to pro- duce high-quality, human-like responses, which brings brand new opportunities for the development of IA paradigms. In this chapter, we identify and introduce two of them in details, i.e., information generation and information synthesis.

generative ai image generation capstone project Diagram
Figure: System Model & Simulation Flow for Generative Ai Image Generation Capstone Project

Information generation allows AI to create tailored content addressing user needs directly, enhancing user experience with immediate, relevant outputs. Information synthesis leverages the ability of generative AI to integrate and reorganize exist- ing information, providing grounded responses and mitigating issues like model hallucination, which is particularly valuable in scenarios requiring precision and external knowledge. This chapter delves into the foundational aspects of gener- ative models, including architecture, scaling, and training, and discusses their applications in multi-modal scenarios. Additionally, it examines the retrieval- augmented generation paradigm and other methods for corpus modeling and understanding, demonstrating how generative AI can enhance information access systems. It also summarizes potential challenges and fruitful directions for future studies.

The primary distinction between modern generative models and traditional AI tech- niques lies in their capability to generate complicated and high-quality output based on human instructions. As shown by many studies , modern generative AI models possess remarkable abilities to generate responses that closely mimic human inter- action. General speaking, such impressive performance comes from their large-scale

Arxiv:2501.02842V1 [Cs.Ir] 6 Jan 2025

training collections and their advanced data modeling algorithms. Their superior data understanding ability can benefit almost every components of existing information access systems, from document encoding and index construction, to query process- ing and relevance analysis, etc. However, when talking about new opportunities or paradigms that are uniquely brought by the generative AI to information access, they can be broadly categorized in two directions. The first one is to create content that directly addresses user’s information needs. By understanding and taking user queries as input instructions, generative AI models are able to generate specific answers or products tailored to the individual’s request. This direct approach to information gen- eration can significantly enhance user experience by providing immediate and relevant responses. The second direction is to leverage the advanced instruction-following capa- bilities of generative AI models to synthesize and recombine existing information in innovative ways. Generative AI such as large language models (LLMs) can take exist- ing data and transform it into new, coherent pieces of information that may not have been explicitly outlined before. This ability to reinterpret and organize information opens up new possibilities for retrieval system design and applications. Therefore, in this chapter, we discuss how generative AI models could help information access from two perspectives, namely information generation and information synthesis.

1 Information Generation

Information need is diverse and typically long-tail. Traditional information retrieval systems, such as search engines and recommendation platforms, are designed to present information that already exists. However, these systems often fall short when it comes to fulfilling the less common information needs. This is particularly evident in scenarios requiring creative creation, where users seek not just information but inspiration and novel ideas. The limitations of traditional information systems in addressing these unique demands have paved the way for the emergence of generative models, which hold the promise of creating new information that aligns closely with the long-tail information needs.

In recent years, generative models have made significant developments. For instance, ChatGPT can respond to user questions, Bing enhances its responses with retrieval-augmented generation, and Midjourney generate images based on user prompts, and recommendation systems generate personal contents for different users.

The development is mainly driven by the capable model architectures, computational resources, and the large-scale internet data. These elements have facilitated the per- formance of generative models to new heights. With the continuous efforts on scaling up these elements, the model performance is still rapidly improving. Nowadays, gener- ative models have gradually been integrated into various workflows and everyday life activities.

In this section, we present the foundation of generative models. This section is organized as follows. Section 1.1 shows the efforts on designing the model architectures for large language models. Section 1.2 discusses how scaling facilitates the development of generative models and its potential future. Section 1.3 presents the different training stages of large language models. Finally, Section 1.4 introduces how large language models are used in multi-modal scenarios.

1.1 Model Architecture

In different generation scenarios like ChatGPT or SoRA, Transformer has emerged as the predominant model structure. It starts with an embedding layer, followed by multiple neural layers. Within each layer, an attention mechanism models the interac- tions between words, creating contextualized embeddings. The final decision on word generation probabilities is derived by comparing the output embedding with the vocab- ulary embeddings. We illustrate the model architecture in Figure 1. Unlike traditional Recurrent Neural Networks , Transformers are capable of modeling long-distance interactions between words directly, which provides a more powerful representational capability. Numerous enhancements to the Transformer architecture have been pro- posed. In the following, we will explore various modifications to each component of the Transformer, highlighting the advancements that have further improved its efficacy and efficiency.

Transformer Layer

Fig. 1 Transformer architecture: the overview on the left and the illustration of one layer on the right .

1.1.1 Word Embedding

Word embedding module is at the bottom of the Transformer architecture. Initially, a tokenizer breaks down a sentence into tokens, which the Word embedding module then maps into embeddings. These are combined with position embeddings and fed into subsequent neural layers. Recent research on large-scale language models has identified word embeddings as one of the main sources to training instability . Particularly in the early stages of training, the gradients of word embeddings are often orders of magnitude larger than those of other parameters. To address this issue, Le Scao et al.

introduced a layer normalization immediately after the word embedding layer, stabilizing the distribution effectively. Besides, Zeng et al. opted to scale down the gradients of the word embeddings by an order of magnitude to prevent substantial updates. Both approaches have been proven effective in stabilizing the training of language models at the 100 billion parameter scale. Yet, whether they are still effective for larger models remains to be investigated.

1.1.2 Position Embedding

Position embedding is essential for Transformer. Unlike RNNs that inherently process sequences in order, vanilla attention mechanism disregards the positional distances between words and Transformer has to rely on position embeddings for position mod- eling. Initially, Transformer utilized Sinusoidal embeddings, a non-trainable form of position embedding that is added directly to word embeddings. Later, Devlin et al.

introduced trainable position embeddings, which is initialized randomly and are updated through gradient descent during training. Subsequently, Raffel et al. and Press et al. proposed relative positioning, where the attention mechanism incorpo- rates biases based on the relative positions of words to better model varying distances.

Recently, Su et al. introduced the concept of rope position embedding, based on the principle that the dot product of vectors correlates with their magnitudes and the angles between them. By rotating vectors in space proportionally to their positions, this method naturally integrates positional information into attention scores. Black et al. has found that this approach outperforms trainable position embeddings.

Yet, these approaches may not work well when extrapolated to long sequences and more effective methods need to be explored.

1.1.3 Attention

The attention mechanism models interactions between words and is a significant com- ponent of the Transformer architecture. Enhancements to the attention module have predominantly focused on two aspects: modeling long texts and optimizing the Key- Value (KV) cache. (1) Modeling Long Texts: The vanilla attention mechanism has a complexity of O(n2), which significantly increases computational costs for long texts. To address this, Sparse Transformer employs sparse attention, utilizing pre- designed attention patterns to avoid the computation of attention over long sequences.

Another approach, Reformer , uses Locality-Sensitive Hashing (LSH) to reduce computational complexity. Additionally, Munkhdalai et al. compressed context information to shorten sequences, thereby reducing overhead. Others have explored retrieval-based methods [16, 17]. This area of research continues to hold considerable potential for future advancements. (2) Optimizing KV Cache: classic Transformers use multi-head attention (MHA), which requires storing extensive key-value caches during inference, slowing down model generation. To mitigate this, Shazeer pro- posed multi-query attention, which employs multiple key heads but only a single value head, substantially reducing the key-value cache and enhancing computational speed. However, Ainslie et al. found that this could degrade model performance, leading to the development of grouped query attention. This method allows multiple key heads to share a single value head, effectively serving as a hybrid between MQA and MHA, balancing computational complexity and performance more effectively.

Recently, DeepSeek-AI introduced multi-head latent attention, which compresses keys and values into a single latent space, thereby reducing the key-value cache while maintaining robust representational capacity.

1.1.4 Layer Normalization

Layer normalization (LayerNorm) is important for stabilizing the distribution of hidden states, a key to train large language models. In the classical Transformer archi- tecture, LayerNorm is positioned between residual blocks, hence termed Post-LN.

Researchers observed that this configuration could lead to high gradients near the output layers and very small gradients near the input layers, resulting in unstable gradients and challenging training dynamics. To address this issue, the Pre-LN configu- ration was proposed , placing LayerNorm on the residual pathways before attention or feed-forward network (FFN) module. Experiments have shown that this adjustment leads to more uniform gradient distribution. Building upon Pre-LN, other researchers introduced Sandwich-LN , which adds an additional LayerNorm at the output of the residual pathways, further enhancing the training stability. Beyond merely adjust- ing the position of LayerNorm, researchers have developed DeepNorm , which combines a tailored parameter initialization strategy with modified residual connec- tions to stabilize training. This approach enables the training of Transformers with depths reaching up to 1000 layers. Nevertheless, there still lacks a theoretical under- standing about how layer normalization affects the training stability and more work needs to be done for scaling the model even further.

1.2 Scaling

Across different information generation scenarios, scaling has been a siginificant factor to the performance improvement. It is largely attributed to the discovery of scaling laws . Scaling laws describe how loss decreases in a log-linear manner as model size or training data volume increases. It can be formulated as follows:

(1)

where L is the loss, x is model size or data size, and k and α are coefficients. This scaling formula has become a crucial theoretical guide in the era of large models, suggesting that performance can be enhanced at a log-linear rate simply by scaling up the model size or training data. Based on these scaling laws, researchers also derived optimal model sizes given fixed computational resources . Their findings indicate that as computational capacity expands, it is beneficial not only to increase the training step but also the model size. This insight has further facilitated the pursuit of large models.

The correctness of scaling laws was first proposed in language modeling field and then validated in many other areas, including data mixture scaling laws , multimodal scaling laws , and scaling laws specific to information retrieval .

Despite wide recognition of scaling laws, there remains disagreement among researchers about whether scaling is the correct path to the future. This stems from two main concerns: the uncertain relationship between loss and practical metrics, and the inference costs associated with large models.

• Loss vs. Metric Improvement: The first arguing point is whether a linear reduction in loss can translate into super-linear improvements in actual metrics. If metrics could improve super-linearly with linear increases in computational effort, scaling

5

up models would be highly advantageous. However, if the decrease in loss only results in linear or sublinear metric improvements, the diminishing improvements make scaling an inefficient option. The relationship between loss and metric per- formance remains an open question. Some researchers believe that metrics can improve super-linearly, which is termed emergent abilities. This is further supported by Du et al. , who observed a jump in metrics when loss reaches a certain thresh- old. Additionally, Power et al. introduced the concept of “grokking” to explain emergence, showing that models might suddenly exhibit strong generalization capa- bilities when provided with sufficient computational resources. Nevertheless, some researchers argued that such phenomena do not exist, showing that a well- trained smaller model can outperform a larger, undertrained one. Schaeffer et al.

demonstrated that emergent abilities are artifacts of discrete metric functions and found that continuous metric functions do not exhibit such behaviors. McKen- zie et al. even found that scaling results in worse metric scores. The existence of specific emergent abilities remains unresolved and needs to be investigated in future work.

• Inference Cost Considerations: Early studies on scaling laws did not account for the higher inference costs associated with larger models. Thus, the arguments that larger models are better do not apply when the inference costs are considered. Instead, small models demonstrate potential to lower the inference costs. As shown by Fang et al. , the optimal model sizes become significantly smaller when accounting for inference costs. Besides, Mei et al. show that smaller models can utilize more sampling steps during inference and thus perform better. Consequently, many recent studies focus on extensively training small models. For example, Llama and MiniCPM are trained with data and steps that far exceed the guidance suggested by scaling laws. In the future, the models may be used on a phone to build up intelligent interaction with users. Thus, it is important to develop high- performing small models.

1.3 Training

Generative models in different scenarios are similar in training. For example, they usually use autoregressive training objectives, pretraining-sft-rlhf training stages, and prompt tuning procedure. In this section, we focus on the text generation scenario. We first discuss the training objectives and then show the three training stages. Finally, we discuss how to design the prompts after the model is trained.

1.3.1 Training Objectives

For generative language models, the training objective is usually next token prediction. However, this was not widely used when Transformers first appeared. Initially, masked language modeling was the prevalent training objective during the BERT era . It masks 15% of the words in a text randomly, and the model is tasked with predicting these masked words. This approach allows the model to utilize bidirectional attention, enhancing its representational capabilities. Even today, BERT models perform bet- ter than autoregressive models on tasks requiring bidirectional attention. However, a

6

significant drawback of this method is the gap between its training setup and down- stream tasks, necessitating a fine-tuning phase for adaptation to various applications. Thus, its zero-shot generalization capabilities are very limited.

Next token prediction was developed to address the inability of masked language modeling to generalize zero-shot to downstream tasks. The authors of GPT-2 proposed that all natural language processing tasks could be reformulated as next token prediction tasks. By training models on this task, models could be directly applied to any downstream task without the need for specific fine-tuning. In fact, research nowadays demonstrates the effectiveness of this idea. Mathematically, next token prediction can be represented with the following formula:

(2)

which is to predict the probability of the next token xt+1 given the sequence of previous tokens.

1.3.2 Training Stages

The training process of language models typically unfolds in three stages: pre-training, supervised fine-tuning (SFT), and reinforcement learning from human feedback (RLHF). Each phase presents unique challenges and methodologies.

Pre-training is the most resource-intensive stage. It is training a randomly ini- tialized model on a large dataset to develop a robust linguistic capability. Several challenges arise during this stage: (1) Large models are especially difficult to train from random initialization. During training, there are often spikes in training loss or difficulty in converging [6, 23, 37]. We discussed various architectural improvements in Section 1.1 to address these instabilities, yet a definitive solution remains an open issue. (2) The computational demand is substantial. Pre-training requires stable and efficient use of computational resources . It often involves parallel processing across multiple machines, which can lead to low utilization rates of computing resources .

Zeng et al. reported numerous hardware failures during pre-training. (3) The qual- ity of pre-training data is crucial . Given the vast amount of data needed, efficiently filtering out low-quality data is essential. The filtering methods usually employ neural scoring models and based on the credibility of the site [40, 41].

Supervised Fine-Tuning (SFT) is to train the model on instruction-response pairs . The model can thus learns to follow instructions or engage in dialogue . To enhance dataset diversity, researchers often leverage different types of NLP tasks. The quality of the dataset is significant and requires a skilled annotation team. Besides, it is also important to label safety-related data, which helps instruct the models to learn to reject inappropriate requests .

Reinforcement Learning from Human Feedback (RLHF) focuses on aligning the model with human preferences based on human feedback [43, 44]. The process starts by sampling real human prompts to which the model generates multiple responses.

These responses are then compared by users or third-party annotators. A reward model is trained based on these human preferences. Subsequently, reinforcement learning

7

techniques utilize the reward model to guide the model updates. This approach sig- nificantly enhances the quality of model outputs, especially in creative writing tasks. However, a major challenge is the generalizability of the reward model; as the model evolves, the reward model may no longer accurately assess the quality of outputs.

Continuous iterations of this process are necessary to mitigate this issue . Recently, there are also some offline reinforcement learning algorithms that do not necessaite training a reward model, such as DPO . Yet studies show that such offline learning methods still underperform the online learning methods.

1.3.3 Prompt Optimization

Generative models are highly sensitive to the input prompts; an effective prompt can significantly enhance the quality of the model’s output . Therefore, optimizing prompts for a generative model is a crucial area of research. Here are three main

Directions:

• Designing Prompt Templates: Researchers often design prompts that mimic human thought processes to guide the model effectively. This includes using structured thought patterns like chain-of-thought , tree-of-thought , and self- consistency , which help the model organize and process information in a logical manner.

• Iterative Optimization of Prompt Templates: like reinforcement learning, this method continuously iterate and refine the prompt templates based on the gen- eration feedback. Given that prompt templates are typically discrete, researchers usually employ large language models to conduct prompt updates [51, 52].

• Training Prompt Rewriting Models Using User Interaction Logs: This approach harnesses the rich feedback contained within user interaction logs to tap into user insights. By analyzing how users interact with the model, researchers can train an automated model to rewrite prompts more effectively. This method leverages real- world data to better align the prompts with user intentions and improve the model’s responses [53, 54].

1.4 Multi-Modal Applications

The rapid advancement of language models has significantly helped progress in the multimodal domain. Language models facilitate the understanding of multimodal data and developments in multimodal generation. We will discuss these two aspects separately.

1.4.1 Multi-Modal Understanding

Multimodal Understanding involves models processing inputs from multiple modali- ties to produce relevant textual responses. For example, GPT-4o can process textual, visual, and auditory input. The challenges in this area include designing model struc- tures that can handle multimodal inputs and crafting appropriate training objectives.

Here, we focus on how visual signals are integrated into large language models: In terms of aligning multimodal inputs, there are mainly three approaches:

8

• Object Detection-Based Input: This method involves detecting objects within an image, extracting their features and associated spatial information, and then feeding this data into the language model [55, 56]. While this approach is effective, it tends to be slow due to the processing time required for object detection.

• Visual Encoding: Another method encodes images directly using a visual encoder, which converts images into a latent vector representation before integration with the model . This method can sometimes result in the loss of detail.

• Patch-Based Input: The most efficient approach involves dividing images into several patches, transforming them with a simple linear layer, and directly inputting them into the model without the need for a complex visual encoder .

In terms of training methods, there are mainly four types of training objectives: • Contrastive Learning or Image-Text matching: These tasks require the model to correctly categorize images and their corresponding textual descriptions, aligning the representations of text and images [61, 63, 64].

• Image Captioning: The model generates captions based on images, which helps it learn to understand the visual content . • Fine-Grained Image Understanding: The model is tasked to describe specific areas of an image or locate particular objects within an image. This helps enhance the model’s detailed comprehension of visual elements [58, 65].

• Image Generation: This task is reconstructing the original pixels of an image that has been blurred or corrupted [58, 66]. These methodologies and training objectives are crucial for advancing models’ capabilities to process and interpret complex multimodal information effectively. This facilitate a more natural interaction with users.

1.4.2 Multi-Modal Generation

Multi-modal generation models, such as text-to-image generation, have substantially revolutionized the field of art creation. Traditionally, GAN and autoregressive methods are mainstream methods. However, they are computationally expensive and can not produce high-quality results. Recently, diffusion [69, 70] emerges as a new state-of-the-art method in multimodal generation. It perturbs the data with noise and learns to reconstruct the original data.

Language models are increasingly applied in the multimodal generation domain, such as in image [71, 72] and video generation [73, 74]. Language models are primarily utilized for processing training data and reformulating prompts.

In terms of training data, the titles associated with real-world images or videos often contain significant noise. If generative models are trained directly on these noisy titles, it could lead to inaccurate semantic understanding. To address this, language models can be used to filter and regenerate text descriptions within the training data [75, 76]. For instance, a multimodal understanding model could first be trained, then used to relabel videos or images to obtain more precise and detailed text descrip- tions. Experimental results have shown that this method significantly improves the fidelity of model generations to prompts.

9

During inference, multimodal generation models are highly sensitive to the input prompts. Many users do not know how to craft effective prompts and thus get unsat- isfying responses . As a result, it is common to train a language model to rewrite user-provided prompts to enhance the quality of the generated images . One of the challenges here is the difficulty in annotating such rewriting training data, as even system developers may not always know the optimal prompts, let alone crowd-sourced workers . To overcome this, some researchers collect a large number of user-shared effective prompts as training data . Others build prompt-rewriting models based on user log data, capturing preferences and feedback for training .

2 Information Synthesis

Other than generating information directly, another important research and appli- cation direction is to use the power of generative AI models, particularly LLMs, to integrate existing information and generate grounded responses accordingly. For sim- plicity, we refer to this paradigm as information synthesis. The key difference between information generation and information synthesis is the source of information. Infor- mation generation relies on the internal knowledge and information gathered through the training of generative AI models to create the model outputs, while informa- tion synthesis requires external sources to provide information to the models, and the models serve more as a integrator than a creator. There are multiple reasons why information synthesis is considered more reliable than generation in several IA sce- narios. Here we discuss two of the most significant ones, i.e., model hallucination and external knowledge.

Hallucinating, which refers to the behavior of generative AI models that create responses and outputs that are not grounded by facts or existing supporting materials, is rooted in the foundation of most existing generative AI systems. For instance, LLMs create responses based on the next token prediction task, which formulates the generation of language as a probabilistic process and generates the next token in the output based on a probabilistic distribution (over the vocabulary) predicted by neural networks [1, 3]. The probabilistic model of LLMs allows them to capture knowledge in large scale data efficiently and effectively, but it also introduces inevitable variance in their generation process. In other words, it is well acknowledged that it’s theoretically impossible to prevent LLMs from generate data that are not seen in their training process . While the ability of hallucinating is the source of creativity for LLMs (and for human as well), it’s not always desirable in practice, particularly for tasks with high requirements on result precision, reliability, and explanability. Therefore, asking the generative AI models to integrate human created or factually grounded materials instead of generating information on their own is often considered more effective and robust to hallucination-sensitive applications.

The need of external knowledge is another key reason why we may prefer informa- tion synthesis over information generation. Despite the fact that modern generative AI models are trained with incredibly large amount of data gathered from the Web, there are many cases where we still need to retrieve and find supports from external

10

knowledge collections to finish certain tasks. Examples including the use of pri- vate datasets, vertical domain applications that require special knowledge, tasks that involve time-sensitive data, etc. It is usually inefficient or prohibitive to update large- scale generative AI models such as LLMs with task-oriented external data through model pre-training or supervised fine-tuning (SFT) . Even if possible, such paradigm is not preferred because the internal knowledge structures of most genera- tive AI models are still mystery (at least of today), and there is no guarantee that the models could behavior and use the external information as we expect. In con- trast, using generative AI models as information synthesizer gives us not only more flexibility, but also more transparency and control over system outputs.

In this section, we discuss how generative AI models, particularly LLMs, can serve as effective information synthesizers for IA. We start with introducing one of the most popular information synthesis paradigm, i.e., retrieval augmented generation (RAG), and then discuss several other directions that utilize LLMs for corpus modeling and understanding.

2.1 Retrieval Augmented Generation

Retrieval Augmented Generation, or RAG, refers to the process of augmenting LLMs with data retrieved from external collections or synthesizing multiple retrieval results with LLMs for downstream applications [84, 85]. While the popularity of RAG rose after the release of large-scale pre-trained language models such as GPT and BART , relevant topics and techniques have already been studied for at least more than two decades in both IR and NLP communities, e.g., extractive and abstractive summarization that generates summary based on retrieved sentences [87, 88] or answer extraction from top retrieved document . A major reason why RAG-like techniques were not as attractive as they are today is the limited performance of generative mod- els before the era of LLMs. After ChatGPT demonstrated superior ability text generation at the end of 2022, there have been many studies and surveys on RAG and its applications in LLMs [84, 90, 91]. As the intent of this chapter is not to provide yet another survey on existing RAG papers, we focus the following discussions on several present and future directions for RAG and their relations underneath.

2.1.1 Naive Rag

Naive RAG refers to the paradigm that directly feeds documents or other types of information retrieved by a retrieval system to the input (e.g., prompts) of a generative AI model and hope that the model can generate better output with or without a specific target task . It is also referred to as the “Retrieve-then-Read” framework that has been used in reading comprehension and text summarization before LLMs hit the world . Given an input (could be a query or a specific task instruction), we first retrieve relevant information (usually entities, passages, or documents) from an external corpus or previous inputs (e.g., the memory of an agent [94, 95]) with a retrieval system. Then, we craft a input prompt with the retrieval results and feed it to the LLM. The LLM will generate the final response based on the input request and

11

the retrieved information. This paradigm has already been proven to be effective in multiple IA tasks such as question answering . Because LLMs are purely used as black-box tools to process the retrieved doc- uments and input request in naive RAG, existing studies on this direction mainly focus on the development of better retrieval systems and prompt design for RAG. The studies on retrieval systems, unsurprisingly, are highly similar to those in IR, which involve indexing, query processing, first-stage retrieval, re-ranking, etc. These topics and system components have already been studied in the IR community for more than five decades. Perhaps the most notable difference is that recent studies on naive RAG often prefer the use of neural retrieval models (e.g., dense retrieval models ) over traditional term-matching models (e.g., BM25 ). An important reason behind this is that neural retrieval models share similar theoretical background and model struc- tures with LLMs. This makes joint optimization possible in modern RAG systems, which we discuss in Section 2.1.3.

The design of input prompts with retrieval results, on the other hand, is relatively more under-explored before the rise of LLMs. It has been well recognized that prompt formats, even when the contents are same, could significantly affect the performance of LLMs. How to feed retrieval results effectively into the prompts of LLMs for RAG has thus attracted a lot of attention recently [93, 98, 99]. Studies have found that LLMs exhibit significant position bias over the input result sequences [100, 101], and has different perspectives on relevance with human experts . Since prompts are the main interaction interface between retrieval and generation, their design principles and downstream effects on naive RAG are of great value both in research and real-world applications. Particularly, how to craft effective RAG prompts automatically could be a fruitful direction to explore. Existing studies have shown that high-quality prompt writers can be automatically learned based on downstream task performance and user logs in image generation , and it is widely believed that similar techniques have also been used in popular LLM chatbots . Yet, how to do this for RAG remains to be a question to be answered.

2.1.2 Modular Rag

In contrast to naive RAG methods, modular RAG treats retrieval systems as func- tional modules to support LLMs . While some works view this retrieval module as one type of many tools that can be learned and used by LLMs , it is widely acknowledged that retrieval systems possesses a irreplaceable position in modern LLM applications due to its diverse nature and significant importance . Broadly speaking, existing studies on using retrieval systems as functional modules for LLM generation mainly focus on the three “W” questions, namely when to retrieve,what to retrieve, and where to retrieve.

The question of when to retrieve refers to the timing of functional call for retrieval systems. In contrast to LLMs that directly create responses based on their internal parameter space without explicit evidence grounding, retrieval systems produces reli- able and explainable information directly by searching external corpus. From this perspective, the best timing to call the retrieval system is when LLMs start to hallu- cinate or produce wrong results. Yet, identifying such timing is difficult because we

12

neither know the correct answers in advance or understand the internal mechanism of LLMs (at least of today) . One naive yet effective method is to retrieve supporting evidence for LLM inference with a fixed time interval, such as every fixed number of generated tokens[107, 108] or every sentence . More advanced paradigms involve the analyze of knowledge boundary and the estimation of prediction uncertainty in LLMs [106, 111]. Theoretically speaking, since the study of when to retrieve shares similar motivations and foundations with the study of hallucination detection, exist- ing studies on LLM hallucination [112, 113] could provide important inspiration for research on this topic. Promising directions including better fact checking systems for LLMs and more investigations on how to characterize the confidence and uncertainty of LLM predictions based on both external behavior and internal state analysis .

The question of what to retrieve focuses on analysis the intents and information needs of LLMs in inference. LLMs often need the help of different tools and systems to finish different tasks . However, in contrast to other tools widely studied in tool learning, retrieval itself is a complicated systems with dynamic and free-form inputs, data collections, and outputs. Therefore, understanding what exactly is needed by LLMs and how to formulate it in the language of retrieval systems is an important problem. Most existing studies on RAG naively use the whole or local context of LLM inference as the queries to retrieval systems and assume that these context contain enough information to guide retrieval . A slightly better solution is to use the terms that LLMs have low confidence to formulate queries since uncertain tokens represent cases where LLMs have limited knowledge to generate responses and thus need more information . As long studied in the IR community, the formulation of an effective query requires deep understanding of the user’s intent, and many of the important context information behind a user intent is not explicitly expressed in the words they wrote . Therefore, a more theoretically principled method to answer what to retrieve in RAG is to analyze the internal state of LLMs and infer their information needs directly. For example, Su et al. directly formulate queries based on the internal attention distribution of LLMs (Figure 2) and improve the performance of RAG for nearly 20% on several benchmark datasets without changing the retrieval system. This demonstrates the potential of future studies on this direction.

Where to retrieve refers to the question of how to identify the correct information sources for RAG. Studies on this direction is particularly related to the research on multi-source retrieval and tool learning . To answer different requests related to the use of information collected from different databases or data collections, LLMs need to learn how to interact with each information sources effectively and efficiently.

The studies of tool learning focus on teaching LLMs to use tools according to the context, and retrieval systems are usually considered as one type of tools to use. However, retrieval itself could be a complicated problem when we possess multiple data collections with different characteristics. In search engines, information sources are broadly categorized based on their modality, and we usually build separate systems for each of them (e.g., the “Images”, “News”, “Videos” tabs on Google). While commercial search engines may aggregate results from different sources into a single page, the ultimate search engine result page (SERP) shown to users are just a list of results and

13

Fig. 2 Su et al. generate queries for RAG based on the internal attention distribution of LLMs. it’s up to the users to decide which they want to see and how to use these results for downstream applications. In contrast, when using LLMs, users often request LLMs to directly answer their question instead of listing a couple of candidates [117, 118], so it’s the job of LLMs to decide where to retrieve the information given the current context. While the studies of how to navigate user queries to search indexes built from different information sources have been widely studied in the IR community [119– 122], how to do it for RAG with modern generative AI models is, to the best of our knowledge, still underexplored. Existing literature on RAG mostly works on a single retrieval collection (usually a text corpus), but it’s obvious that no single collection can satisfy the needs of LLMs in different tasks. For instance, when writing a legal

14

case document, the judge needs to collect and organize information from evidences, complaints, counterclaims, court records, as well as legal articles and previous cases. How to navigate the generation model to retrieve and integrate information from different sources jointly for downstream applications is a practical and potentially fruitful research question for RAG.

2.1.3 Optimization Of Retrieval And Generation

As discussed in several RAG surveys [84, 90], the optimization of RAG systems usually involves the optimization of three components, i.e., the retriever, the generator, and the augmentation method. If we further step back and look at the high-level goals of RAG optimization, we could also categorize it based on how we evaluate the RAG system, namely the evaluation from the perspectives of retrievers, generators, or the joint systems. The evaluation from the retriever perspectives is not particularly different from existing studies on ranking evaluation. The underlining assumption of this is that, once the LLMs are fed with the passages or documents that contain the correct information, they should be able to produce the correct answers directly. Therefore, the evaluation and optimization of a RAG system could downgrade to the evaluation and optimization of a classic retrieval/ranking systems, to where most existing works on dense retrieval and LTR could be applied [123, 124]. Yet, there are still differences between RAG and traditional retrieval tasks as the queries are no long issued by users.

How to formulate queries efficiently and effectively from LLMs for the retriever is worthy research question, and studies on this direction has already shown potentials in improving the overall quality of RAG systems .

From the perspective of generators, RAG evaluation and optimization focus more on improving the robustness and effectiveness of LLM generation based on a fixed set of retrieval results . This often means extra training or fine-tuning on LLMs to improve their fundamental ability in information processing. For example, retrieved documents could be lengthy, and LLMs are usually not good at processing long input context . Therefore, how to design efficient LLMs that can take long context inputs efficiently and effectively has been a popular research problem that have been widely studied by researchers from both academia and industry . We have seen many companies show off their models based on how many input tokens they can process in one request. In addition, since retrieval results are fed as a part of the LLM inputs, whether the LLMs can generate the response based on the retrieved documents instead of their internal knowledge could be seen as a special type of instruction-following ability. Studies have been conducted to teach LLMs to utilize retrieval results faithfully and constantly in RAG systems On the other hand, factors such as irrelevant results and ranking perturbations are well acknowledged to be harmful for the performance of generators in RAG, so there are also studies that try to improve the robustness of LLMs from the perspective of RAG. For example, Zhang et al. proposes to fine tune LLMs with the presence of retrieval results (i.e., retrieval augmented fine tuning) so that LLMs can learn the domain-specific knowledge introduced by the retriever and improve their robustness against potential distracting information from retrieval.

15

From the perspective of augmentation methods, existing research mostly focuses on the joint optimization of RAG system as a whole. In other words, the loss functions of RAG optimization should be built from the performance metrics of downstream tasks directly. While this paradigm is appealing, it often has strict requirements on the design of RAG systems. Particularly, it’s difficult to apply such joint optimiza- tion algorithms on a RAG system in which retrievers and generators are loosely connected through prompts constructed from discrete retrieval results. While rein- forcement learning could solve the problem in theory, its empirical performance when being used as the solo optimization algorithms for ranking systems is still not satisfying at this point . If you already have a good retriever and only conduct fine-tuning with a fixed LLM, then it may work , but this still doesn’t look like a perfect solution because reinforcement learning usually subject to large variance in practice.

To the best of our knowledge, how to directly connect the training of retrievers with the auto-regressive loss of the generators in RAG is still an open question. Answering this question requires us to go deep into the structure of generative AI models and retrieval models, and develop new model structures that can take advantages from studies on both sides.

2.1.4 Retrieval Planning and Composite Information Needs As discussed above, the initial motivation behind the studies of RAG mostly focuses on using the power of retrieval systems to improve the quality of responses generated by LLMs in terms of reliability and informativeness. While it is widely acknowledged that problems such as hallucination and high computation cost in supervised fine- tuning will continue to be significant for generative AI models in a short period of time, there are also concerns, especially from the IR community, that retrieval could become less important with the rapid evolution of LLMs . In fact, ChatGPT has already shown similar accuracy and better user satisfaction on factoid question answering than traditional web search engines . However, the rise of generative AI models also brings brand new opportunities for IR. One of them is the possibility of moving from SERPs that simply list result candidates to a real information agent that solve complicated tasks with composite information needs.

Today, most people treat IR systems as unit information solvers. Despite of their actual task characteristics, users first decompose their goals into a couple of unit information need (usually expressed with separate queries), and then issue them one by one to search engines or recommendation systems to find the corresponding answers.

An important reason behind the popularity of this paradigm is that, at least of today, IR systems are not capable of doing complicated information tasks with composite needs and multi-step planning. For example, we can use a search engine to find a survey on RAG by searching ”survey of RAG”, but cannot write such a survey directly by retrieving and analyzing papers from publication collections. The job of information need decomposition and retrieval planning has always been human’s.

Fortunately, with the help of generative AI models like LLMs, it is now possi- ble to push the boundary of IR systems and tackle such advanced information tasks for users. Composite retrieval is not a new concept in IR , but previous studies refer to the phrase as retrieval paradigms that cluster results from multiple sources

16

and show them in groups for specific user queries . While this represents one type of composite needs, it is relatively simple as the target user queries usually are mostly topic-specific and keyword-based. Complicated information tasks such as sur- vey generation and professional document writing often involve multi-step planning and multi-round interactions between the retrieval results and response generation. To build powerful IR systems or agents that can solve such composite information tasks, we need to construct collaborative systems that deeply connect the retrieval, plan- ning, and generation. For instance, we need to conduct generation-oriented retrieval optimization to build retrieval framework and model interfaces for downstream task planner and response generators; we also need to design retrieval-oriented generation models that can decompose information needs, navigate the retrieval process, gather information from multiple sources to generate the final results. Research on these direc- tions could be fruitful and significantly extend the scope of IR in the era of generative

2.2 Corpus Modeling And Understanding

In contrast to using RAG, another line of studies try to use generative AI models to replace traditional retrieval systems. Directly answering user’s information need instead of showing ten blue links has long been an important goal for the development of intelligent IR systems . With the rise of LLMs, such vision is now achievable in a significant extent. For example, LLM-based chatbots like ChatGPT can answer multiple types of user queries with direct answers . Metzler et al. has dis- cussed several paradigms in which pre-trained language models can help IR systems answer user’s information needs directly without listing references. The intuition is to use neural network based language models to store the corpus knowledge in parameter space and pull relevant answers or information directly from it based on user’s queries.

Depending on how the problem is formulated, several research directions have emerged. Specifically, in this section, we discuss two of them, namely generative retrieval and domain-specific modeling.

2.2.1 Generative Retrieval

The idea of Generative Retrieval comes from the idea of differentiable index proposed by Metzler et al. . The original name used in the paper was Model-based IR, but after the rise of generative AI models, some researchers start to refer to studies on this direction as generative retrieval (GR). The core idea of GR is two-fold, i.e., the differentiable index and the generation of doc IDs.

Inspired by the superior performance of pre-trained language models, particularly BERT and GPT , generative retrieval wants to explore the possibility of replacing traditional term-based index (e.g., inverted index) in retrieval systems with large- scale neural networks. In contrast to dense retrieval models that build neural encoders to project documents to latent semantic spaces and build explicit indexes based on document vectors, GR tries to build implicit indexes in the parameter space of neural networks. For instance, DSI and its variations have tried to train pretrained language models on the target corpus directly and then treat the model’s parameter

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.