AdaViT: Adaptive Tokens for Efficient Vision Transformer

Hongxu Yin, Arash Vahdat, Jose Alvarez, Arun Mallya, Jan Kautz, Pavlo Molchanov

Introduction

Transformers have emerged as a popular class of neural network architecture that computes network outputs using highly expressive attention mechanisms. Originated from the natural language processing (NLP) community, they have been shown effective in solving a wide range of problems in NLP, such as machine translation, representation learning, and question answering . Recently, vision transformers have gained an increasing popularity in the vision community and they have been successfully applied to a broad range of vision applications, such as image classification , object detection , image generation , and semantic segmentation . The most popular paradigm remains when vision transformers form tokens via splitting an image into a series of ordered patches and perform inter-/intra-calculations between tokens to solve the underlying task. Processing an image with vision transformers remains computationally expensive, primarily due to the quadratic number of interactions between tokens . Therefore, deploying vision transformers on data processing clusters or edge devices is challenging amid significant computational and memory resources.

The main focus of this paper is to study how to automatically adjust the compute in visions transformers as a function of the complexity of the input image. Almost all mainstream vision transformers have a fixed cost during inference that is independent from the input. However, the difficulty of a prediction task varies with the complexity of the input image. For example, classifying a car versus a human from a single image with a homogeneous background is relatively simple; while differentiating between different breeds of dogs on a complex background is more challenging. Even within a single image, the patches that contain detailed object features are far more informative compared to those from the background. Inspired by this, we develop a framework that adaptively adjusts the compute used in vision transformers based on the input.

The problem of input-dependent inference for neural networks has been studied in prior work. Graves proposed adaptive computation time (ACT) to represent the output of the neural module as a mean-field model defined by a halting distribution. Such formulation relaxes the discrete halting problem to a continuous optimization problem that minimizes an upper bound on the total compute. Recently, stochastic methods were also applied to solve this problem, leveraging geometric-modelling of exit distribution to enable early halting of network layers . Figurnov et al. proposed a spatial extension of ACT that halts convolutional operations along the spatial cells rather than the residual layers. This approach does not lead to faster inference as high-performance hardware still relies on dense computations. However, we show that the vision transformer’s uniform shape and tokenization enable an adaptive computation method to yield a direct speedup on off-the-shelf hardware, surpassing prior work in efficiency-accuracy tradeoff.

In this paper, we propose an input-dependent adaptive inference mechanism for vision transformers. A naive approach is to follow ACT, where the computation is halted for all tokens in a residual layer simultaneously. We observe that this approach reduces the compute by a small margin with an undesirable accuracy loss. To resolve this, we propose A-ViT, a spatially adaptive inference mechanism that halts the compute of different tokens at different depths, reserving compute for only discriminative tokens in a dynamic manner. Unlike point-wise ACT within convolutional feature maps , our spatial halting is directly supported by high-performance hardware since the halted tokens can be efficiently removed from the underlying computation. Moreover, entire halting mechanism can be learnt using existing parameters within the model, without introducing any extra parameters. We also propose a novel approach to target different computational budgets by enforcing a distributional prior on the halting probability. We empirically observe that the depth of the compute is highly correlated with the object semantics, indicating that our model can ignore less relevant background information (see quick examples in Fig. A-ViT: Adaptive Tokens for Efficient Vision Transformer and more examples in Fig. 3). Our proposed approach significantly cuts down the inference cost – A-ViT improves the throughput of DEIT-Tiny by 62%62\% and DEIT-Small by 38%38\% with only 0.3%0.3\% accuracy drop on ImageNet1K.

We introduce a method for input-dependent inference in vision transformers that allows us to halt the computation for different tokens at different depth.

We base learning of adaptive token halting on the existent embedding dimensions in the original architecture and do not require extra parameters or compute for halting.

We introduce distributional prior regularization to guide halting towards a specific distribution and average token depth that stabelizes ACT training.

We analyze the depth of varying tokens across different images and provide insights into the attention mechanism of vision transformer.

We empirically show that the proposed method improves throughput by up to 62%62\% on hardware with minor drop in accuracy.

Related Work

There are a number of ways to improve the efficiency of transformers including weight sharing across transformer blocks , dynamically controlling the attention span of each token , allowing the model to output the result in an earlier transformer block , and applying pruning . A number of methods have aimed at reducing the computationally complexity of transformers by reducing the quadratic interactions between tokens . We focus on approaches related to adaptive inference that depends on the input image complexity. A more detailed analysis of the literature is present in .

Special architectures. One way is to change the architecture of the model to support adaptive computations . For example, models that represent a neural network as a fixed-point function can have the property of adaptive computation by default. Such models compute the difference to the internal state and, when applied over multiple iterations, converge towards the solution (desired output). For example, neural ordinary differential equations (ODEs) use a new architecture with repetitive computation to learn the dynamics of the process . Using ODEs requires a specific solver, is often slower than fix depth models and requires adding extra constraints on the model design. learns a set of classifiers with different resolutions executed in order; computation stops when confidence of the model is above the threshold. proposed a residual variant with shared weights and a halting mechanism.

Stochastic and reinforcement learning (RL) methods. The depth of a residual neural network can be changed during inference by skipping a subset of residual layers. This is possible since residual networks have the same input and output feature dimensions and they are known to perform feature refinements iteratively. Individual extra models can be learned on the top of a backbone to change the computational graph. A number of approaches proposed to train a separate network via RL to decide when to halt. These approaches require training of a dedicated halting model and their training is challenging due to the high-variance training signal in RL. Conv-AIG learns conditional gating of residual blocks via Gumbel-softmax trick. extends the idea to spatial dimension (pixel level).

Adaptive inference in vision transformers. With the increased popularity, researchers have very recently explored adaptive inference for vision transformers. DynamicViT uses extra control gates that are trained with the Gumbel-softmax trick to halt tokens and it resembles some similarities to Conv-AIG and . Gumbel-softmax-based relaxation solutions might be sub-optimal due to the difficulty of regularization, stochasticity of training, and early convergence of the stochastic loss, requiring multi-stage token sparsification as a heuristic guidance. In this work, we approach the problem from a rather different perspective, and we study how an ACT -like approach can be defined for spatially adaptive computation in vision transformers. We show complete viability to remove the need for the extra halting sub-networks, and we show that our models bring simultaneous efficiency, accuracy, and token-importance allocation improvements, as shown later.

A-ViT

Consider a vision transformer network that takes an image x∈RC×H×Wx\in\mathcal{R}^{C\times H\times W} (CC, HH, and WW represent channel, height, and width respectively) as input to make a prediction through:

where the encoding network E(⋅)\mathcal{E}(\cdot) tokenizes the image patches from xx into the positioned tokens t∈RK×Et\in\mathcal{R}^{K\times E}, KK being the total number of tokens and EE the embedding dimension of each token. C(⋅)\mathcal{C}(\cdot) post-processes the transformed class token after the entire stack, while the LL intermediate transformer blocks F(⋅)\mathcal{F}(\cdot) transform the input via self-attention. Consider the transformer block at layer ll that transforms all tokens from layer l−1l-1 via:

where t1:Klt_{1:K}^{l} denotes all the KK updated token, with t1:K0=E(x)t_{1:K}^{0}=\mathcal{E}(x). Note that the internal computation flow of transformer blocks F(⋅)\mathcal{F}(\cdot) is such that the number of tokens KK can be changed from a layer to another. This offers out-of-the-box computational gains when tokens are dropped due to the halting mechanism. Vision transformer utilizes a consistent feature dimension EE for all tokens throughout layers. This makes it easy to learn and capture a global halting mechanism that monitors all layers in a joint manner. This also makes halting design easier for transformers compared to CNNs that require explicit handling of varying architectural dimensions, e.g., number of channel, at different depths.

To halt tokens adaptively, we introduce an input-dependent halting score for each token as a halting probability hklh_{k}^{l} for a token kk at layer ll:

where H(⋅)H(\cdot) is a halting module. Akin to ACT , we enforce the halting score of each token hklh^{l}_{k} to be in the range 0≤hkl≤10\leq h^{l}_{k}\leq 1, and use accumulative importance to halt tokens as inference progresses into deeper layers. To this end, we conduct the token stopping when the cumulative halting score exceeds 1−ϵ1-\epsilon:

where ϵ\epsilon is a small positive constant that allows halting after one layer. To further alleviate any dependency on dynamically halted tokens between adjacent layers, we mask out a token tkt_{k} for all remaining depth l>Nkl>N_{k} once it is halted by (i) zeroing out the token value, and (ii) blocking its attention to other tokens, shielding its impact to tl>Nkt^{l>N_{k}} in Eqn. 2. We define h1:KL=1h_{1:K}^{L}=\mathbf{1} to enforce stopping at the final layer for all tokens. Our token masking keeps the computational cost of our training iterations similar to the original vision transformer’s training cost. However, at the inference time, we simply remove the halted tokens from computation to measure the actual speedup gained by our halting mechanism.

We incorporate H(⋅)H(\cdot) into the existing vision transformer block by allocating a single neuron in the MLP layer to do the task. Therefore, we do not introduce any additional learnable parameters or compute for halting mechanism. More specifically, we observe that the embedding dimension EE of each token spares sufficient capacity to accommodate learning of adaptive halting, enabling halting score calculation as:

where tk,elt^{l}_{k,e} indicates the ethe^{\text{th}} dimension of token tklt_{k}^{l} and σ(u)=11+exp−u\sigma(u)=\frac{1}{1+\text{exp}^{-u}} is the logistic sigmoid function. Above, β\beta and γ\gamma are shifting and scaling parameters that adjust the embedding before applying the non-linearity. Note that these two scalar parameters are shared across all layers for all tokens. Only one entry of the embedding dimension EE is used for halting score calculation. Empirically, we observe that the simple choice of e=0e=0 (the first dimension) performs well, while varying indices does not change the original performance, as we show later. As a result our halting mechanism does not introduce additional parameters or sub-network beyond the two scalar parameters β\beta and γ\gamma.

To track progress of halting probabilities across layers, we calculate a remainder for each token as:

that subsequent forms a halting probability as:

Given the range of hh and rr, halting probability per token at each layer is always bounded 0≤pkl≤10\leq p^{l}_{k}\leq 1. The overall ponder loss to encourage early stopping is formulated via auxiliary variable rr (reminder):

where ponder loss ρk\rho_{k} of each token is averaged. Vision transformers use a special class token tkt_{k} to produce the classification prediction, we denote it as tct_{c} for future notations. This token similar to other input tokens is updated in all layers. We apply a mean-field formulation (halting-probability weighted average of previous states) to form the output token tot_{o} and the associated task loss as:

Our vision transformer can then be trained by minimizing:

where αp\alpha_{\text{p}} scales the pondering loss relative to the the main task loss. Algorithm 1 describes the overall computation flow, and Fig. 2 depicts the associated halting mechanism for visual explanation. At this stage, the objective function encourages an accuracy-efficiency trade-off when pondering different tokens at varying depths, enabling adaptive control.

One critical factor in Eqn. 10 is αp\alpha_{\text{p}} that balances halting strength and network performance for the target application. A larger αp\alpha_{\text{p}} value imposes a stronger penalty, and hence learns to halt tokens earlier. Despite efficacy towards computation reduction, prior work on adaptive computation have found that training can be sensitive to the choice of αp\alpha_{\text{p}} and its value may not provide a fine-grain control over accuracy-efficiency trade-off. We empirically observe a similar behavior in vision transformers.

As a remedy, we introduce a distributional prior to regularize hlh^{l} such that tokens are expected to exit at a target depth on average, however, we still allow per-image variations. In this case for infinite number of input images we expect the the depth of token to vary within the distributional prior. Similar prior distribution has been recently shown effective to stablize convergence during stochastic pondering . To this end, we define a halting score distribution:

that averages expected halting score for all tokens across at each layer of network (i.e., H∈RL\mathcal{H}\in\mathcal{R}^{L}). Using this as an estimate of how halting likelihoods distribute across layers, we regularize this distribution towards a pre-defined prior using KL divergence. We form the new distributional prior regularization term as:

where KL refers to the Kullback-Leibler divergence, and Htarget\mathcal{H}^{\text{target}} denotes a target halting score distribution with a guiding stopping layer. We use the probability density function of Gaussian distribution to define a bell-shaped distribution Htarget\mathcal{H}^{\text{target}} in this paper, centered at the expected stopping depth NtargetN^{\text{target}}. Intuitively, this weakly encourages the expected sum of halting score for each token to trigger exit condition at NtargetN^{\text{target}}. This offers enhanced control of expected remaining compute, as we show later in experiments.

Our final loss function that trains the network parameters for adaptive token computation is formulated as:

where αd\alpha_{\text{d}} is a scalar coefficient that balances the distribution regularization against other loss terms.

Experiments

We evaluate our method for the classification task on the large-scale 10001000-class ImageNet ILSVRC 2012 dataset at the 224×224224\times 224 pixel resolution. We first analyze the performance of adaptive tokens, both qualitatively and quantitatively. Then, we show the benefits of the proposed method over prior art, followed by a demonstration of direct throughput improvements of vision transformers on legacy hardware. Finally, we evaluate the different components of our proposed approach to validate our design choices.

Implementation details. We base A-ViT on the data-efficient vision transformer architecture (DeiT) that includes 1212 layers in total. Based on original training recipeBased on official repository at https://github.com/facebookresearch/DeiT., we train all models on only ImageNet1K dataset without auxiliary images. We use the default 16×1616\times 16 patch resolution. For all experiments in this section, we use Adam for optimization (learning rate 1.5⋅10−31.5\cdot 10^{-3}) with cosine learning rate decay. For regularization constants we utilize αd=0.1,αp=5⋅10−4\alpha_{\text{d}}=0.1,\alpha_{\text{p}}=5\cdot 10^{-4} to scale loss terms. We use γ=5,β=−10\gamma=5,\beta=-10 for sigmoid control gates H(⋅)H(\cdot), shared across all layers. We use the embedding value at index e=0e=0 to represent the halting probability (H(⋅)H(\cdot)) for tokens. Starting from publicly available pretrained checkpoints, we fine-tune DeiT-T/S variant models for 100100 epochs, respectively, to learn adaptive tokens without distillation. We denote the associated adaptive versions as A-ViT-T/S respectively. In what follows, we mainly use the A-ViT-T for ablations and analysis before showing efficiency improvements for both variants afterwards. We find that mixup is not compatible with adaptive inference, and we focus on classification without auxiliary distillation token – we remove both from finetuning. Applying our finetuning on the full DeiT-S and DeiT-T results in a top-1 accuracy of 78.9%78.9\% and 71.3%71.3\%, respectively. For training runs we use 88 NVIDIA V100 GPUs and automatic-mixed precision (AMP) acceleration.

Qualitative results. Fig. 3 visualizes the tokens’ depth that is adaptively controlled during inference with our A-ViT-T over the ImageNet1K validation set. Remarkably, we observe that our adaptive token halting enables longer processing for highly discriminative and salient regions, often associated with the target class. Also, we observe a highly effective halting of relatively irrelevant tokens and their associated computations. For example, our approach on animal classes retains the eyes, textures, and colors from the target object and analyze them in full depth, while using fewer layers to process the background (e.g., the sky around the bird, and ocean around sea animals). Note that even background tokens marked as not important still actively participate in classification during initial layers. In addition, we also observe the inspiring fact that adaptive tokens can easily (i) keep track of repeating target objects, as shown in the first image of the last row in Fig. 3, and (ii) even shield irrelevant objects completely (see second image of last row).

Token depth distribution. Given a complete distinct token distribution per image, we next analyze the dataset-level token importance distributions for additional insights. Fig. 4 (a) depicts the average depth of the learnt tokens over the validation set. It demonstrates a 2D Gaussian-like distribution that is centered at the image center. This is consistent with the fact that most ImageNet samples are centered, intuitively aligning with the image distribution. As a result, more compute is allocated on-the-fly to center areas, and computational cost on the sides is reduced.

Halting score distribution. To further evaluate the halting behavior across transformer layers, we plot the average layer-wise halting score distribution over 1212 layers. Fig. 4 (b) shows box plots of halting scores averaged over all tokens per layer per image. The analysis is performed on 55K randomly sampled validation images. As expected, the halting score gradually increases at initial stages, peaks and then decreases for deeper layers.

Sharp-halting baseline. To further compare with static models of the same depth for performance gauging, we also train a DeiT-T with 88 layers as a sharp-halting baseline. We observe that our A-ViT-T outperforms this new baseline by +1.4%+1.4\% top-1 accuracy at a similar throughput. Although our adaptive regime is on average similarly shallow, it still inherits the expressivity of the original deeper network, as we observe that informative tokens are processed by deeper layers (e.g., until 12th12^{\text{th}} layer as in Fig. 3).

Easy and hard samples. We can analyse the difficulty of an image for the network by looking at the averaged depth of the adaptive tokens per image. Therefore, in Fig. 5, we depict hard and easy samples in terms of the required computation. Note, all samples in the figure are correctly classified, and only differ by the averaged token depth. We can observe that images with homogeneous background are relatively easy for classification, and A-ViT processes them much faster than hard samples. Hard samples represent images with informative visual features distributed over the entire image, and hence incur more computation.

Class-wise sensitivity. Given an adaptive inference paradigm, we analyze the change in classification accuracy for various classes with respect to the full model. In particular, we compute class-wise validation accuracy changes before and after applying adaptive inference. We summarize both qualitative and quantitative results in Table 1. We observe that originally very confident or uncertain samples are not affected by adaptive inference. Adaptive inference improves accuracy of the visually dominant classes such as individual furniture and animals.

2 Comparison to Prior Art

Next, we compare our method with previous work that study adaptive computation. For comprehensiveness, we systematically compare with five state-of-the-art halting mechanisms, covering both vision and NLP methods that tackle the dynamic inference problem from different perspectives: (i) adaptive computation time as ACT reference applied on halting entire layers, (ii) confidence-based halting that gauges on logits, (iii) similarity-based halting that oversees layer-wise similarity, (iv) pondering-based halting that exits based on stochastic halting-probabilities, and (v) the very recent DynamicViT that learns halting decisions via Gumble-softmax relaxation. Details in appendix.

Performance comparison. We compare our results in Table 2 and demonstrate simultaneous performance improvements over prior art in having smaller averaged depth, smaller number of FLOPs and better classification accuracy. Notably our method involves no extra parameters, while cutting down FLOPs by 39%39\% with only a minor loss of accuracy. To further visualize improvements over the state-of-the-art DynamicViT , we include Fig. 6 as a qualitative comparison of token depth for an official sample presented in the work. As noticed, A-ViT more effectively captures the important regions associated with the target objects, ignores the background tokens, and improves efficiency.

Note that both DynamicViT and A-ViT investigate adaptive tokens but from two different angles. DynamicViT utilizes Gumbel-Softmax to learn halting and incorporates a control for computation via a multi-stage token keeping ratio; it provides stronger guarantees on the latency by simply setting the ratio. A-ViT on the other hand takes a complete probabilistic approach to learn halting via ACT. This enables it to freely adjust computation, and hence capture enhanced semantic and improve accuracy, however requires a distributional prior and has a less intuitive hyper-parameter.

Hardware speedup. In Table 3, we compare speedup on off-the-shelf GPUs. See appendix for measurement details. In contrast to spatial ACT in CNNs that require extra computation flow and kernel re-writing , A-ViT enables speedups out of the box in vision transformers. With only 0.3%0.3\% in accuracy drop, our method directly improves the throughputs of DeiT small and tiny variants by 38%38\% and 62%62\% without requiring hardware/library modification.

3 Ablations

Here, we perform ablations studies to evaluate each component in our method and validate their contributions.

Token-level ACT via Lponder\mathcal{L}_{\text{ponder}}. One noticeable distinction of this work from conventional ACT is a full exploration of spatial redundancy in image patches, and hence their tokens. Comparing the first and last row in Table 2, we observe that our fine-grained pondering reduces token depths by roughly 33 layers, and results in 25%25\% more FLOP reductions compared to the conventional ACT.

Distributional prior via Ldistr.\mathcal{L}_{\text{distr.}}. Incorporating the distributional prior allows us to better guide the expected token depth towards a target average depth, as seen in Fig. 7. As opposed to αp\alpha_{\text{p}} that indirectly gauges on the remaining efficiency and usually suffers from over-/under-penalization, our distributional prior guides a quick convergence to a target depth level, and hence improves final accuracy. Note that a distributional prior complements the ponder loss in guiding overall halting towards a target depth, but it cannot capture remainder information – using ACT-agnostic distributional prior alone results in an accuracy drop of more than 2%2\%.

“Free” embedding to learn halting. Next we justify the usage of a single value in the embedding vector for halting score computation and representation. In the embedding vectors, we set one entry at a random index to zero and analyze the associated accuracy drop without any finetuning of the model. Repeating 1010 times for DeiT-T/S variants, the ImageNet1K top-1 accuracy only drops by 0.08%±0.04%0.08\%\pm 0.04\%/0.04%±0.03%0.04\%\pm 0.03\%, respectively. This experiment demonstrates that one element in the vector can be used for another task with minimal impact on the original performance. In our experiments, we pick the first element in the vector and use it for the halting score computation.

Layer-wise networks to learn halting. We continue to examine viability to leverage extra networks for halting learning. To this end we add an extra two-layer learnable network (with input/hidden dimensions of 192/96192/96, internal/output gates as GeLU/Sigmoid) on top of embeddings of each layer in A-ViT-T. We observed a very slight increase in accuracy of +0.06%+0.06\% with +0.2+0.2M parameter and −12.6%-12.6\% inference throughput overhead, as auxiliary nets have to be executed sequentially with ViT layers. Given this tradeoff, we base learning of halting on existing ViT parameters.

Limitations & Future Directions

In this work we primarily focused on the classification task. However, extension to other tasks such as video processing can be of great interest, given not only spatial but also temporal redundancy within input tokens.

Conclusions

We have introduced A-ViT to adaptively adjust the amount of token computation based on input complexity. We demonstrated that the method improves vision transformer throughput on hardware without imposing extra parameters or modifications of transformer blocks, outperforming prior dynamic approaches. Captured token importance distribution adaptively varies by input images, yet coincides surprisingly well with human perception, offering insights for future work to improve vision transformer efficiency.

References

Appendix B - Additional Details

For training setup other than the scaling constants, lr specified in the main manuscript, we follow original repository for all other hyper-parameters at https://github.com/facebookresearch/DeiT such as drop out rate, momentum, preprocessing, etc, imposing minimum training recipe changes when adapting a static model to its adaptive counterpart.

Latency

We measure the latency on an NVIDIA TITAN RTX 20802080 GPU with PyTorch for batch size of 6464 images, CUDA 10.210.2. For GPU warming up, 100100 forward passes are conducted, and then the median speed of the 11K measurements of the full model latency are reported. The exact same setup is shared across all baseline and proposed methods for a fair comparison.

SOTA baselines.

We followed DeiT’s repositoryhttps://github.com/facebookresearch/DeiT for recipes and checkpoints as a common starting point for all experiments. For DynamicViT , we used the public repository and script from the authors. For other dynamic approaches from CNN/NLP literature, we re-implemented the methods on DeiT to examine ACT for layer-wise halting, confidence threshold on post-softmax logits, two variants of similarity gauging on delta-logits based on (i) LPIPS and (ii) MSE similarity scores, and PonderNet with geometric-distribution sampling towards token halting. For all methods, a detailed grid search was conducted to ensure optimal hyper-parameters.