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

Rag Retrieval Augmented 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

Prompt

(c) Patch-based Autoregressive Retrieval Augmentation (Ours) ...

Retrieved Images

...

?

Figure 1: Comparison between Autoregressive Retrieval Augmentation (AR-RAG) for image generation in (c) and existing image generation paradigms in (a) (b). In AR-RAG, image patches in red boxes denote retrieval queries and keys, image patches in blue boxes are retrieved values, and gray boxes with the question mark are next image patches to be predicted. (Caption: A white cat is

Abstract

We introduce Autoregressive Retrieval Augmentation (AR-RAG), a novel paradigm that enhances image generation by autoregressively incorporating k- nearest neighbor retrievals at the patch level. Unlike prior methods that perform a single, static retrieval before generation and condition the entire generation on fixed reference images, AR-RAG performs context-aware retrievals at each gen- eration step, using prior-generated patches as queries to retrieve and incorporate the most relevant patch-level visual references, enabling the model to respond to evolving generation needs while avoiding limitations (e.g., over-copying, stylis- tic bias, etc.) prevalent in existing methods. To realize AR-RAG, we propose two parallel frameworks: (1) Distribution-Augmentation in Decoding (DAiD), a training-free plug-and-use decoding strategy that directly merges the distribu- tion of model-predicted patches with the distribution of retrieved patches, and (2) Feature-Augmentation in Decoding (FAiD), a parameter-efficient fine-tuning method that progressively smooths the features of retrieved patches via multi-scale convolution operations and leverages them to augment the image generation pro- cess. We validate the effectiveness of AR-RAG on widely adopted benchmarks, 1Jingyuan Qi and Zhiyang Xu contributed equally to this work.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Preprint. Under review. including Midjourney-30K, GenEval and DPG-Bench, demonstrating significant performance gains over state-of-the-art image generation models.1

Introduction

Recent advancements in image generation have demonstrated remarkable capabilities in producing photorealistic images based on user prompts [31, 28, 7, 37, 10, 41, 9, 43, 45, 27, 6]. However, despite these improvements, the generated images often exhibit local distortions and inconsistencies, particularly in visual objects that possess complex structures , frequently interact with other objects and the surrounding scene [22, 26], or are underrepresented in the training data . A promising approach to mitigating these challenges is retrieval-augmented generation (RAG), which enhances the generation process by incorporating real-world images as additional references [8, 3].

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

While RAG has been extensively explored in the language domain [23, 13], its application to image and multimodal generation remains largely underdeveloped. A few existing studies [3, 8, 46, 48, 49] bridge this gap by performing a single-step retrieval based on the input prompt prior to generation, conditioning the entire image generation process on fixed visual cues (Figure 1 (b)). However, as demonstrated in our pilot study (Section 5.2), such static, coarse-grained retrieval approaches [3, 8, 49] frequently introduce irrelevant or weakly aligned visual contents that persist throughout generation.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Since the retrieved images are selected once, before decoding begins, and remain unchanged, these methods cannot respond to the evolving generation needs, resulting in over-copying of irrelevant details, stylistic bias, and the hallucination of unrelated visual elements. For example, as shown in Figure 1(b), a basketball player present in the retrieved references, despite being irrelevant to the input prompt, unintentionally appears in the generated image.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

In this paper, we propose Autoregressive Retrieval Augmentation (AR-RAG), a novel retrieval- augmented paradigm for image generation that dynamically and autoregressively incorporates patch- level k-nearest-neighbor (k-NN) retrievals throughout the generation process (Figure 1(c)). In contrast to prior methods that rely on static, coarse-grained retrievals of entire reference images, typically using captions as retrieval queries and keys, AR-RAG performs fine-grained, step-wise retrieval at the image patch level. Specifically, as generation unfolds, AR-RAG leverages the already- generated surrounding patches as localized queries to retrieve contextually similar patches from a pre-constructed patch-level database. This database is built by encoding real-world images into latent patch features, where each entry contains a patch embedding as a value and the embeddings of its h-hop spatial neighbors as a key. During the generation of the next target patch (gray boxes in Figure 1(c)), AR-RAG retrieves the top-K most relevant patches (blue boxes) by measuring similarity between the surrounding generated context patches (red boxes) and database keys (also red boxes). These retrieved patches are then integrated into the model to inform and enhance the prediction of the next patch, enabling the model to dynamically adjust to local generation needs.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

By conditioning on the evolving generation context as retrieval queries, AR-RAG ensures that retrieved visual references remain relevant throughout the generation process, encouraging local semantic coherence. Moreover, the patch-level retrieval allows for precise integration of visual elements without overcommitting to entire reference images, avoiding the limitations of over-copying or irrelevant conditioning observed in static retrieval.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

To realize the AR-RAG framework, we introduce two parallel implementations: (1) Distribution- Augmentation in Decoding (DAiD), a training-free, plug-and-play decoding strategy that merges the model’s predicted patch distribution with that of the retrieved patches. Specifically, the top-K retrieved patches are assigned probabilities inversely proportional to their normalized ℓ2 distances computed from the query and key patch embeddings. These probabilities are then linearly combined with the model’s native output distribution to guide the next patch prediction, enabling retrieval-aware generation without any additional training. (2) Feature-Augmentation in Decoding (FAiD), a parameter-efficient fine-tuning approach that integrates retrieved patches into the generation process through learned smoothing and blending mechanisms. Specifically, when generating the next image token, FAiD operates in two stages: (1) refining the retrieved patch features by adjusting them to better fit the local context of the already generated surrounding patches, based on parameterized convolutional operations of varying kernel sizes; and (2) blending the refined features of retrieved patches with the model’s predicted feature representation for the next patch, based on compatibility 1Code and model checkpoints can be found at https://github.com/PLUM-Lab/AR-RAG.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

2

scores computed for each retrieved patch to quantify their alignment with the current generation context. To enable iterative refinement, we insert multiple FAiD modules at selected transformer layers, where the output of each FAiD module, i.e., the context-aware retrieved features blended at that layer, is forwarded as input to the next FAiD module in deeper layers. This progressive retrieval refinement mechanism allows the model to incrementally enhance its predictions as patch- level representations evolve through the network. We evaluate AR-RAG on three widely adopted benchmarks, including Midjourney-30K 2, Geneval , and DPG-Bench . Experimental results demonstrate that both DAiD and FAiD significantly improve the coherence and naturalness of generated images while introducing only marginal computational overhead.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

The contributions of our work can be summarized as follows: • We propose AR-RAG, the first patch-level autoregressive retrieval augmentation framework which dynamically retrieves and integrates fine-grained visual content to enhance image generation, while avoiding limitations (e.g., over-copying, stylistic bias, etc.) prevalent in existing image-level retrieval augmentation methods.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

• We introduce Distribution-Augmentation in Decoding (DAiD), a training-free, plug-and-play decoding strategy that directly integrates the distribution of retrieved patches into that predicted by the image generation models, enabling easy integration into existing architectures.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

• We introduce Feature-Augmentation in Decoding (FAiD), a parameter-efficient fine-tuning frame- work that progressively refines and blends retrieval signals via lightweight convolutional modules, enhancing spatial coherence and visual quality across layers.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

• Extensive experiments and analysis show that AR-RAG significantly improves performance of state-of-the-art image generation model across diverse metrics. In particular, Janus-Pro with FAiD achieves 6.67 FID on Midjourney-30K and 0.78 overall score on GenEval, establishing a new state of the art among autoregressive image generation models of comparable scale.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

We Implement Both Daid And Faid Based On Janus-

Pro , an autoregressive (AR) unified generation model, due to its strong performance. Janus-Pro is initialized from a transformer-based pre-trained large-language model , and employs a quantized autoencoder to encode images into discrete image tokens. During multimodal pretraining, the model learns to predict a sequence of discrete image tokens [v1, v2, ...vN] conditioned on an input text prompt [t1, t2, ...tM]. The training objective is formally defined as:

(1)

where D is the training corpus. This is the same training objective used in our FAiD method in Section 3.3. We argue that DAiD and FAiD can be extended to any image generation model that autoregressively predicts probability distributions of discrete image tokens such as LlamaGen , Show-o and VAR .

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Quantized Autoencoder

The quantized autoencoder used in Janus-Pro consists of an encoder θenc, a decoder θdec, and a codebook Z. The encoder, a convolutional neural network, downsamples and compresses raw pixel inputs into compact patch representations. During the quantization process, each patch representation is mapped to an index in the codebook by identifying its nearest neighbor vector in the codebook. In the decoding stage, these patch indices are mapped back to their corresponding vector representations via the codebook, and the decoder, another convolutional neural network, reconstructs the image from these compact representations. In our implementation, we leverage this autoencoder to build the coupled database for Janus-pro which is detailed in Section 3.1.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

3

AR-RAG: Patch-based Autoregressive Retrieval Augmentation

Patch-Based Retrieval Database Construction

We build a patch-based retrieval database based on several large-scale, real-world image datasets, including CC12M and JourneyDB . Specifically, for each image I, we encode it into N 2https://huggingface.co/datasets/playgroundai/MJHQ-30K

Dretrieval

Figure 2: The decoding process in Distribution-Augmentation in Decoding (DAiD). patches using the quantized autoencoder , θEnc, from Janus-Pro: V = θenc(I) ∈R

N×D,

where d is the hidden dimension, and Vij corresponds to the latent representation of the patch at position (i, j). We utilize each patch vector Vij as the value of a database entry and the representation of its h-hop surrounding patches as the key. Here, the h-hop surrounding patch representation is formed by concatenating the vectors of adjacent patches centering around (i, j) in a top-to-bottom, left-to-right order. For example, for a patch at position (i, j), the 1-hop surrounding representation spans 8 surrounding patches [V(i−1)(j−1) : V(i−1)(j) : V(i−1)(j+1) : V(i)(j−1) : V(i)(j+1) : V(i+1)(j−1) : V(i+1)(j) : V(i+1)(j+1)] where : denotes the concatenation operation of image patch features. If a patch is located at the edge of the image and lacks certain surrounding patches, we substitute each missing surrounding patch with a zero vector 0.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Distribution-Augmentation In Decoding (Daid)

Given a text prompt T, Janus-Pro autoregressively predicts a sequence of image tokens [v1, v2, ...vN] where per-token probability is defined in Equation 1. As shown in Figure 2, DAiD augments this process by incorporating probability distributions from retrieved image patches. Specifically, when Janus-Pro predicts the next image token vij, we first utilize the codebook Z to convert vij’s h- hop already generated surrounding patches into patch representations. If no surrounding image tokens are available at a given position (e.g., when i = 0 or j = 0), we use the zero vector 0 as a placeholder. Once we compute the representation of vij’s h-hop surrounding patches, we leverage it as the retrieval query and retrieve the top-K most similar patch representations from the database constructed in Section 3.1 using l2 distance. We denote the representations of the top-K retrieved patches as [ˆv1, ˆv2, ..., ˆvK] and their corresponding l2 distances as [s1, s2, ..., sK]. These retrieved representations are then mapped back to discrete token indices using the codebook: ˆvk = Z(ˆvk).

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

To augment the generation process with the retrieved image tokens [ˆv1, ˆv2, ..., ˆvK], we create a retrieval-based distribution Dretrieval ∈R|Z| over the entire codebook Z, where |Z| is the codebook size. Tokens not included in the top-K retrieved set are assigned a probability of 0. For tokens within the top-K, we compute their probabilities using a softmax over their l2 distance to the query, scaled

(3)

This creates a sparse distribution where only the top-K retrieved tokens have non-zero probabilities. Finally, we merge this retrieval distribution with the model’s predicted distribution Dmodel using a

(4)

where λ ∈[0, 1] is the retrieval weight hyperparameter controlling the influence of retrieved patches on the final distribution. The next token is then sampled from this merged distribution: vij ∼Dmerge.

Feature-Augmentation In Decoding (Faid)

While DAiD offers a training free approach to directly augment the probability distribution of predicted patches using retrieved ones, it suffers from noise propagation and limited flexibility in

Retrieval Database

...

Spb

...

10

Figure 3: Overall architecture of Feature-Augmentation in Decoding (FAiD). fully leveraging the fine-grained visual information in the retrieved patches. We thus further propose FAiD, a feature-based autoregressive augmentation strategy to enhance the image generation process.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

As illustrated in Figure 3, when predicting the next token vij during image generation, we employ the same retrieval process described in Section 3.2 to obtain the top-K most relevant patches and their representations [ˆv1, ˆv2, ..., ˆvK] from our database. To effectively incorporate them into the autoregressive generation process, FAiD consists of two steps: (1) refining retrieved patches to ensure coherence with the surrounding context of vij in the generated image, and (2) adaptively blending the representation of refined patches with the hidden state of the predicted next patch based on learned compatibility scores. To enable progressive refinement of retrieved information as representations evolve through the network, we insert a FAiD module for every L/b decoder layers of the generation model, where L denotes the total number of decoder layers and b is a hyperparameter.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Multi-Scale Feature Smoothing

The key of effective patch integration lies in ensuring spatial coherence between retrieved patches and the surrounding image context. To achieve this, we propose multi-scale feature smoothing (Algorithm 1 in Appendix A), where multi-scale convolutions are applied to retrieved patches within the generation context, so that the retrieved visual features are smoothed to preserve structural and stylistic consistency with the surrounding context of the predicted token. Specifically, at each step when predicting the next image token vij, we first construct a 2D

√

N×D of the current partially-generated image by arranging

N−1, Hl

n] from the current decoder layer l. We use 0 vectors as placeholders for positions that have not yet been generated. Then, we transform the retrieved patch representations [ˆv1, ˆv2, ..., ˆvK] into the generation model’s hidden space by mapping each patch ˆvk to a discrete token index via the codebook Z and embedding it through the pretrained image

Embedding Layer Embimg:

[ˆh1, ˆh2, ..., ˆhK] = Embimg([Z(ˆv1), Z(ˆv2), ..., Z(ˆvK)])

(5)

For each retrieved patch ˆhk, we create a copy of Hl where position (i, j) (the location of vij) is replaced with ˆhk. We then apply convolution operations at multiple scales (2 × 2 through Q × Q) to capture contextual patterns at different resolutions. To maintain computational efficiency, we only perform convolution operations when the kernel covers position (i, j), rather than processing the entire image. Each convolution kernel Convq×q produces a refined representation ˆhq

K For The Retrieved

patch at scale q. The final refined representation for each retrieved patch is computed as a weighted

(6)

where Ω= [ω2, ..., ωQ] are learnable parameters that determine the importance of each scale.

Feature Augmentation

After feature smoothing, some of the retrieved patch features may still not be able to fit into the surrounding neighbors and hence we need to lower their impact in the final repre- sentation. Thus, we compute a compatibility score for each of the refined patches. This is achieved by projecting each refined retrieved patch representation through a linear transformation parameterized by a weight matrix W ∈R1×D, yielding the score sk = ˆhkWT . The final representation for the next image token vij after layer j is computed as:

Ij Is The Residual, ∆Hl

ij is the updated representation from the transformer layer l, and

Pk

k=1 skˆhk is the contribution of the retrieved image patches.

Patch-Based Retrieval Database

To construct our patch-level retrieval database, we randomly sample 5.7 million images from CC12M , 3.3 million from JourneyDB , and 4.6 million from DataComp , while ensuring that any samples included in the testing set are excluded to prevent data leakage. Each image is encoded into a sequence of patch-level representations and image tokens using the same image tokenizer employed in the Janus-Pro model. For efficient similarity search, we implement our retriever using the FAISS library .

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Training Setup

We adopt Janus-Pro-1B and Show-o as our backbone models and fine-tune them on a dataset of 50,000 image-caption pairs sampled from CC12M and Midjourney-v6 3. We empirically determine the optimal hyperparameters for DAiD and FAiD, and the complete hyperparameter optimization experiment results can be found in Appendix C.2. Further details regarding the training dataset construction and implementation can be found in Appendix B.3.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Baselines

To evaluate the effectiveness of our proposed methods, we adopt several state-of-the-art image generation approaches as baselines, including non-retrieval models such as LlamaGen , LDM , Stable Diffusion (SDv1.5 and SDv3) [31, 10], PixArt-alpha , DALL-E 2 , Show- o, and Janus-Pro , and image-based retrieval augmentation methods, including RDM , RA-CM3 , and ImageRAG. Since pretrained models of RA-CM3 are not publicly available, we try our best to replicate their method based on Janus-Pro to ensure a fair comparison. More details of training and implementation of RA-CM3 can be found in Appendix B.1.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Evaluation Benchmarks And Metrics

To comprehensively evaluate our proposed methods, we employ three benchmarks: (1) GenEval , which assesses models’ ability to generate images with specific attributes and relationships described in text prompts; (2) DPG-Bench , which evaluates performance on detailed prompts with complex requirements; and (3) Midjourney-30k , where we employ three complementary metrics: FID for measuring statistical similarity between generated and real image distributions, CMMD for assessing alignment with human perception using CLIP embeddings, and FWD for evaluating spatial and frequency coherence through wavelet packet coefficients. For all three metrics, lower scores indicate higher quality generated images. Detailed descriptions of these benchmarks and metrics can be found in Appendix B.4.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Text-To-Image Generation Results

Tables 1, 2, and 3 present performance comparisons across multiple benchmarks, where our AR- RAG methods consistently outperform existing approaches. Notably, previous retrieval-augmented approaches such as RDM and ImageRAG perform worse than their non-retrieval counterparts (LDM and SDXL, respectively) on both GenEval and DPG-Bench. We provide detailed analysis for existing image-level retrieval methods and highlight the unique advantages of our AR-RAG frameworks in the following discussion and Section 5.2. Appendix C.1 provides a benchmark analysis to demonstrate the effectiveness of patch-level retrieval in our AR-RAG methods.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

3https://huggingface.co/datasets/brivangl/midjourney-v6-llava

Params

Single Obj. Two Obj.

Position

Color Attri.

0.78 (+0.07)

Table 1: Evaluation of text-to-image generation ability on GenEval benchmark. Note our methods are based on Janus-Pro highlighted in gray.

79.36 (+2.10)

Table 2: Evaluation of text-to-image generation ability on DPG-Bench. Note our methods are based on Janus-Pro highlighted in gray. On GenEval, our methods show significant improvements in categories such as “Two Obj.” and “Position,” which demand accurate multi-object generation and spatial arrangement. These gains are largely due to the local and dynamic nature of our autoregressive patch-level retrieval. Consider the prompt “a green couch and an orange umbrella”, a combination that rarely co-occurs in real-world images. Static full-image retrieval methods may retrieve references containing only one of the objects.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Taking these references as a global visual prior throughout the generation can lead the model to overfit to irrelevant layouts or dominant visual structures in the retrieved examples. On DPG-Bench, which features dense and highly detailed prompts, the performance gap between our method and prior retrieval-augmented approaches becomes even more substantial. Similar as GenEval, existing

Table 3: Evaluation Of Text-To-Image Generation

ability on the Midjourney-30K benchmark.

Image-Level Retrieval Augmentation Methods Strug-

gle to retrieve meaningful references when the num- ber of distinct entities and attributes in a prompt increases.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Mentation Framework Overcomes This Limitation By

dynamically retrieving patch-level visual features

The Original Prompt, Enabling More Targeted And

effective augmentation.

Sistently Outperform Both Janus-Pro And Show-O

baselines across all three evaluation metrics. No- tably, despite operating locally at the patch level, our approach leads to a significant reduction in FID

7

scores, indicating improved global visual quality and closer alignment with the distribution of real images. This suggests that context-aware, auto-regressive retrieval and refinement can propagate to enhance holistic image fidelity. Furthermore, the improvements in CMMD and FWD metrics confirm our method’s effectiveness in reducing visual distortions and enhancing coherence. These results also demonstrate that AR-RAG delivers robust and architecture-agnostic improvements, validating its broad applicability across different image generation backbones.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

A Photo Of A

bench.

Of Taylor Swift

with a red scarf.

Glow On A Pair Of

high-top sneakers.

Beside A Plush,

round red couch.

Faid

A photo of a sheep. Figure 4: Qualitative results of DAiD, FAiD and baselines.

Bench (Left Three Columns) And

GenEval (right two columns).

Gans), And Implausible Configu-

rations (e.g., column 4, where a chair exhibits an impossible design). Both DAiD and FAiD substantially reduce such local distortions, with FAiD yielding the highest visual quality. These results confirm that autoregressive retrieval effectively maintains object consistency and structural integrity throughout the generation process, particularly for complex objects and multi-object scenes.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

(d) A photo of a green couch and an orange umbrella. (c) A photo of a green cup and a yellow bowl.

Imagerag

(a) A photo of an apple.

Ar-Rag

(b) A photo of a white dog and a blue potted plant.

Imagerag

Figure 5: Images generated by ImageRAG and our AR-RAG. ImageRAG excessively copies retrieved images and does not follow user prompts. Figure 5 presents a comparative analysis of conventional image-level and our autoregressive patch- level retrieval augmentation methods. By comprehensively examining images produced by Im- ageRAG alongside their corresponding retrieved reference images, we identify two critical challenges inherent in image-level retrieval augmentation approaches. First, these methods tend to overcopy irrelevant visual elements from retrieved reference images into the generation outputs. As illustrated in Figure 5 (a), when generating an image of an apple, image-level retrieval approaches retrieve a reference image showing an apple on a tree branch and subsequently incorporate both the apple and the surrounding branches, despite the prompt making no mention of them. Similarly, for the prompt “a green cup and a yellow bowl” in Figure 5 (b), the image-level retrieval augmentation approach retrieves a green Starbucks cup and reproduces the pattern on the cup in the generated image,

8

despite this element not being part of the original instruction. This overcopying behavior directly compromises the instruction-following capability of generative models. Figure 5 (c) demonstrates that when prompted to generate “A photo of a white dog and a blue potted plant,” image-level retrieval methods produce an image containing only the white dog, omitting the blue potted plant entirely.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Similarly, for “a photo of a green couch and an orange umbrella” in Figure 5 (d), the generated image fails to include the umbrella. This degradation in instruction following occurs because image-level retrieval biases the generation process toward the compositional structure of retrieved reference images, which may not align with the multi-object relationships specified in the prompt. In contrast, by autoregressively retrieving and integrating visual information at the fine-grained patch level rather than the image level, AR-RAG enables selective incorporation of relevant visual elements while maintaining independence from irrelevant contextual features present in the reference images.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Table 4: Inference Time For Generat-

ing 100 images on a single L40 card. Table 4 shows the inference time comparisons across different models when generating 100 images using both a single L40 GPU. The DAiD method introduces only a minimal increase in inference time compared to the base Janus-Pro-1B model, with an average overhead of just 0.22%, demonstrating that DAiD maintains high computational efficiency. FAiD shows a more noticeable overhead of 36.03% on a single GPU due to its autoregressive retrieval and feature blending operations.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

However, this increase remains reasonable given the substan- tial performance gains in generation quality. Overall, both DAiD and FAiD do not significantly compromise the inference efficiency of Janus-Pro, making them practical for real-world applications.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Related Work

Retrieval-augmented generation (RAG) has emerged as a powerful paradigm that enhances generative models by incorporating external knowledge during decoding [23, 13, 16, 47, 46, 15, 24, 25, 42]. Originally developed for natural language processing, RAG enables models to retrieve relevant documents to supplement parametric knowledge during response generation , and has been widely adopted in many downstream tasks, such as knowledge-intensive tasks , document fusion , model pretraining , dialogue generation [35, 1], and so on.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Beyond the text domain, prior research has explored enhancing image generation by incorporating external visual references. Early approaches [8, 3] condition the diffusion process on retrieved images, typically encoded via CLIP or VAE encoders, to guide generation toward higher visual fidelity. KNN-Diffusion extends this idea by leveraging k-nearest neighbor images to improve zero-shot generalization to novel domains. Building on this retrieval-augmented framework, more recent methods [49, 33] introduce adaptive retrieval pipelines that iteratively refine retrieved images based on feedback from multimodal large language models (MLLMs) analyzing the generated outputs.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

These methods enable context-aware and prompt-sensitive guidance during generation. Another line of work encodes multimodal retrievals into discrete visual and text tokens, and uses them directly as contextual input to augment the generation process of a multimodal large language model. All of these works differe from our method by that our method works on patch-level, enabling more fine grain retrievals and can dynamically adjust retrievals based on evolving generation states.

rag retrieval augmented generation capstone project Diagram
Figure: System Model & Simulation Flow for Rag Retrieval Augmented Generation Capstone Project

Conclusion

In this work, we propose Autoregressive Retrieval Augmentation (AR-RAG), a novel retrieval paradigm that enhances image synthesis by leveraging k-nearest neighbor retrievals at the patch level. Unlike traditional image-level retrieval approaches, AR-RAG enables fine-grained visual element integration while maintaining compositional flexibility. We introduce two parallel frameworks: (1) Distribution-Augmentation in Decoding (DAiD), a training-free approach that integrates retrieved patch distributions directly into generation, and (2) Feature-Augmentation in Decoding (FAiD), which employs parameter-efficient fine-tuning with multi-scale feature smoothing and compatibility-based feature augmentation. Extensive experiments across GenEval, DPG-Bench, and Midjourney-30K demonstrate that AR-RAG significantly outperforms both conventional and retrieval-augmented baselines, particularly in handling complex prompts with multiple objects and specific spatial rela-

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.