EdgeNeXt: Efficiently Amalgamated CNN-Transformer Architecture for Mobile Vision Applications

Muhammad Maaz, Abdelrahman Shaker, Hisham Cholakkal, Salman Khan, Syed Waqas Zamir, Rao Muhammad Anwer, Fahad Shahbaz Khan

Introduction

Convolutional neural networks (CNNs) and the recently introduced vision transformers (ViTs) have significantly advanced the state-of-the-art in several mainstream computer vision tasks, including object recognition, detection and segmentation . The general trend is to make the network architectures more deeper and sophisticated in the pursuit of ever-increasing accuracy. While striving for higher accuracy, most existing CNN and ViT-based architectures ignore the aspect of computational efficiency (i.e., model size and speed) which is crucial to operating on resource-constrained devices such as mobile platforms. In many real-world applications e.g., robotics and self-driving cars, the recognition process is desired to be both accurate and have low latency on resource-constrained mobile platforms.

Most existing approaches typically utilize carefully designed efficient variants of convolutions to achieve a trade-off between speed and accuracy on resource-constrained mobile platforms . Other than these approaches, few existing works employ hardware-aware neural architecture search (NAS) to build low latency accurate models for mobile devices. While being easy to train and efficient in encoding local image details, these aforementioned light-weight CNNs do not explicitly model global interactions between pixels.

The introduction of self-attention in vision transformers (ViTs) has made it possible to explicitly model this global interaction, however, this typically comes at the cost of slow inference because of the self-attention computation . This becomes an important challenge for designing a lightweight ViT variant for mobile vision applications.

The majority of the existing works employ CNN-based designs in developing efficient models. However, the convolution operation in CNNs inherits two main limitations: First, it has local receptive field and thereby unable to model global context; Second, the learned weights are stationary at inference times, making CNNs inflexible to adapt to the input content. While both of these issues can be alleviated with Transformers, they are typically compute intensive. Few recent works have investigated designing lightweight architectures for mobile vision tasks by combining the strengths of CNNs and ViTs. However, these approaches mainly focus on optimizing the parameters and incur higher multiply-adds (MAdds) operations which restricts high-speed inference on mobile devices. The MAdds are higher since the complexity of the attention block is quadratic with respect to the input size . This becomes further problematic due to multiple attention blocks in the network architecture. Here, we argue that the model size, parameters, and MAdds are all desired to be small with respect to the resource-constrained devices when designing a unified mobile architecture that effectively combines the complementary advantages of CNNs and ViTs (see Fig. 1).

Contributions. We propose a new light-weight architecture, named EdgeNeXt, that is efficient in terms of model size, parameters and MAdds, while being superior in accuracy on mobile vision tasks. Specifically, we introduce split depth-wise transpose attention (SDTA) encoder that effectively learns both local and global representations to address the issue of limited receptive fields in CNNs without increasing the number of parameters and MAdd operations. Our proposed architecture shows favorable performance in terms of both accuracy and latency compared to state-of-the-art mobile networks on various tasks including image classification, object detection, and semantic segmentation. Our EdgeNeXt backbone with 5.6M parameters and 1.3G MAdds achieves 79.4% top-1 ImageNet-1K classification accuracy which is superior to its recently introduced MobileViT counterpart , while requiring 35% less MAdds. For object detection and semantic segmentation tasks, the proposed EdgeNeXt achieves higher mAP and mIOU with fewer MAdds and a comparable number of parameters, compared to all the published lightweight models in literature.

Related Work

In recent years, designing lightweight hardware-efficient convolutional neural networks for mobile vision tasks has been well studied in literature. The current methods focus on designing efficient versions of convolutions for low-powered edge devices . Among these methods, MobileNet is the most widely used architecture which employs depth-wise separable convolutions . On the other hand, ShuffleNet uses channel shuffling and low-cost group convolutions. MobileNetV2 introduces inverted residual block with linear bottleneck, achieving promising performance on various vision tasks. ESPNetv2 utilizes depth-wise dilated convolutions to increase the receptive field of the network without increasing the network complexity. The hardware-aware neural architecture search (NAS) has also been explored to find a better trade-off between speed and accuracy on mobile devices . Although these CCNs are faster to train and infer on mobile devices, they lack global interaction between pixels which limits their accuracy.

Recently, Desovitskiy et al. introduces a vision transformer architecture based on the self-attention mechanism for vision tasks. Their proposed architecture utilizes large-scale pre-training data (e.g. JFT-300M), extensive data augmentations, and a longer training schedule to achieve competitive performance. Later, DeiT proposes to integrate distillation token in this architecture and only employ training on ImageNet-1K dataset. Since then, several variants of ViTs and hybrid architectures are proposed in the literature, adding image-specific inductive bias to ViTs for obtaining improved performance on different vision tasks .

ViT models achieve competitive results for several visual recognition tasks . However, it is difficult to deploy these models on resource-constrained edge devices because of the high computational cost of the multi-headed self-attention (MHA). There has been recent work on designing lightweight hybrid networks for mobile vision tasks that combine the advantages of CNNs and transformers. MobileFormer employs parallel branches of MobileNetV2 and ViTs with a bridge connecting both branches for local-global interaction. Mehta et al. consider transformers as convolution and propose a MobileViT block for local-global image context fusion. Their approach achieves superior performance on image classification surpassing previous light-weight CNNs and ViTs using a similar parameter budget.

Although MobileViT mainly focuses on optimizing parameters and latency, MHA is still the main efficiency bottleneck in this model, especially for the number of MAdds and the inference time on edge devices. The complexity of MHA in MobileViT is quadratic with respect to the input size, which is the main efficiency bottleneck given their existing nine attention blocks in MobileViT-S model. In this work, we strive to design a new light-weight architecture for mobile devices that is efficient in terms of both parameters and MAdds, while being superior in accuracy on mobile vision tasks. Our proposed architecture, EdgeNeXt, is built on the recently introduced CNN method, ConvNeXt , which modernizes the ResNet architecture following the ViT design choices. Within our EdgeNeXt, we introduce an SDTA block that combines depth-wise convolutions with adaptive kernel sizes along with transpose attention in an efficient manner, obtaining an optimal accuracy-speed trade-off.

EdgeNeXt

The main objective of this work is to develop a lightweight hybrid design that effectively fuses the merits of ViTs and CNNs for low-powered edge devices. The computational overhead in ViTs (e.g., MobileViT ) is mainly due to the self-attention operation. In contrast to MobileViT, the attention block in our model has linear complexity with respect to the input spatial dimension of O(Nd2)\mathcal{O}(Nd^{2}), where NN is the number of patches, and dd is the feature/channel dimension. The self-attention operation in our model is applied across channel dimensions instead of the spatial dimension. Furthermore, we demonstrate that with a much lower number of attention blocks (3 versus 9 in MobileViT), we can surpass their performance mark. In this way, the proposed framework can model global representations with a limited number of MAdds which is a fundamental criterion to ensure low-latency inference on edge devices. To motivate our proposed architecture, we present two desirable properties.

a) Encoding the global information efficiently. The intrinsic characteristic of self-attention to learn global representations is crucial for vision tasks. To inherit this advantage efficiently, we use cross-covariance attention to incorporate the attention operation across the feature channel dimension instead of the spatial dimension within a relatively small number of network blocks. This reduces the complexity of the original self-attention operation from quadratic to linear in terms of number of tokens and implicitly encodes the global information effectively.

b) Adaptive kernel sizes. Large-kernel convolutions are known to be computationally expensive since the number of parameters and FLOPs quadratically increases as the kernel size grows. Although a larger kernel size is helpful to increase the receptive field, using such large kernels across the whole network hierarchy is expensive and sub-optimal. We propose an adaptive kernel sizes mechanism to reduce this complexity and capture different levels of features in the network. Inspired by the hierarchy of the CNNs, we use smaller kernels at the early stages, while larger kernels at the latter stages in the convolution encoder blocks. This design choice is optimal as early stages in CNN usually capture low-level features and smaller kernels are suitable for this purpose. However, in later stages of the network, large convolutional kernels are required to capture high-level features . We explain our architectural details next.

Overall Architecture. Fig. 2 illustrates an overview of the proposed EdgeNeXt architecture. The main ingredients are two-fold: (1) adaptive N×NN{\times}N Conv. encoder, and (2) split depth-wise transpose attention (SDTA) encoder. Our EdgeNeXt architecture builds on the design principles of ConvNeXt and extracts hierarchical features at four different scales across the four stages. The input image of size H×W×3H{\times}W{\times}3 is passed through a patchify stem layer at the beginning of the network, implemented using a 4×44{\times}4 non-overlapping convolution followed by a layer norm, which results in H4×W4×C1\frac{H}{4}{\times}\frac{W}{4}{\times}C1 feature maps. Then, the output is passed to 3×\times3 Conv. encoder to extract local features. The second stage begins with a downsampling layer implemented using 2×{\times}2 strided convolution that reduces the spatial sizes by half and increases the channels, resulting in H8×W8×C2\frac{H}{8}{\times}\frac{W}{8}{\times}C2 feature maps, followed by two consecutive 5×{\times}5 Conv. encoders. Positional Encoding (PE) is also added before the SDTA block in the second stage only. We observe that PE is sensitive for dense prediction tasks (e.g., object detection and segmentation) as well as adding it in all stages increases the latency of the network. Hence, we add it only once in the network to encode the spatial location information. The output feature maps are further passed to the third and fourth stages, to generate H16×W16×C3\frac{H}{16}{\times}\frac{W}{16}{\times}C3 and H32×W32×C4\frac{H}{32}{\times}\frac{W}{32}{\times}C4 dimensional features, respectively.

Convolution Encoder. This block consists of depth-wise separable convolution with adaptive kernel sizes. We can define it by two separate layers: (1) depth-wise convolution with adaptive N×NN{\times}N kernels. We use kk = 3, 5, 7, and 9 for stages 1, 2, 3, and 4, respectively. Then, (2) two point-wise convolution layers are used to enrich the local representation alongside standard Layer Normalization (LN) and Gaussian Error Linear Unit (GELU) activation for non-linear feature mapping. Finally, a skip connection is added to make information flow across the network hierarchy. This block is similar to the ConvNeXt block but the kernel sizes are dynamic and vary depending on the stage. We observe that adaptive kernel sizes in Conv. encoder perform better compared to static kernel sizes (Table 8). The Conv. encoder can be represented as follows:

where xi\bm{x}_{i} denotes the input feature maps of shape H<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><molspace="0em"rspace="0em">×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord"><spanclass="mord">×</span></span></span></span></span></span>W<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><molspace="0em"rspace="0em">×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord"><spanclass="mord">×</span></span></span></span></span></span>CH<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo lspace="0em" rspace="0em">×</mo></mrow><annotation encoding="application/x-tex">{\times}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord"><span class="mord">×</span></span></span></span></span></span>W<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo lspace="0em" rspace="0em">×</mo></mrow><annotation encoding="application/x-tex">{\times}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord"><span class="mord">×</span></span></span></span></span></span>C, LinearGLinear_{G} is a point-wise convolution layer followed by GELU, DwDw is k<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><mo>×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord">×</span></span></span></span></span>kk<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo>×</mo></mrow><annotation encoding="application/x-tex">\times</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord">×</span></span></span></span></span>k depth-wise convolution, LNLN is a normalization layer, and xi+1\bm{x}_{i+1} denotes the output feature maps of the Conv. encoder.

SDTA Encoder. There are two main components in the proposed split depth-wise transpose attention (SDTA) encoder. The first component strives to learn an adaptive multi-scale feature representation by encoding various spatial levels within the input image and the second part implicitly encodes global image representations. The first part of our encoder is inspired by Res2Net where we adopt a multi-scale processing approach by developing hierarchical representation into a single block. This makes the spatial receptive field of the output feature representation more flexible and adaptive. Different from Res2Net, the first block in our SDTA encoder does not use the 1×11{\times 1} pointwise convolution layers to ensure a lightweight network with a constrained number of parameters and MAdds. Also, we use adaptive number of subsets per stage to allow effective and flexible feature encoding. In our STDA encoder, we split the input tensor H<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><molspace="0em"rspace="0em">×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord"><spanclass="mord">×</span></span></span></span></span></span>W<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><molspace="0em"rspace="0em">×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord"><spanclass="mord">×</span></span></span></span></span></span>CH<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo lspace="0em" rspace="0em">×</mo></mrow><annotation encoding="application/x-tex">{\times}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord"><span class="mord">×</span></span></span></span></span></span>W<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo lspace="0em" rspace="0em">×</mo></mrow><annotation encoding="application/x-tex">{\times}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord"><span class="mord">×</span></span></span></span></span></span>C into ss subsets, each subset is denoted by xi\bm{x}_{i} and has the same spatial size with C/sC/s channels, where ii ∈\in {1, 2, …, ss} and CC is the number of channels. Each feature maps subset (except the first subset) is passed to 3×33{\times}3 depth-wise convolution, denoted by did_{i}, and the output is denoted by yi\bm{y}_{i}. Also, the output of di−1d_{i-1}, denoted by yi−1\bm{y}_{i-1}, is added to the feature subset xi\bm{x}_{i}, and then fed into did_{i}. The number of subsets ss is adaptive based on the stage number tt, where tt ∈\in {2, 3, 4}. We can write yi\bm{y}_{i} as follows:

Each depth-wise operation did_{i}, as shown in SDTA encoder in Fig. 2, receives feature maps output from all previous splits {xj\bm{x}_{j}, jj ≤\leq ii}.

As mentioned earlier, the overhead of the transformer self-attention layer is infeasible for vision tasks on edge-devices because it comes at the cost of higher MAdds and latency. To alleviate this issue and encode the global context efficiently, we use transposed query and key attention feature maps in our SDTA encoder . This operation has a linear complexity by applying the dot-product operation of the MSA across channel dimensions instead of the spatial dimension, which allows us to compute cross-covariance across channels to generate attention feature maps that have implicit knowledge about the global representations. Given a normalized tensor Y\bm{Y} of shape H<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><molspace="0em"rspace="0em">×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord"><spanclass="mord">×</span></span></span></span></span></span>W<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><molspace="0em"rspace="0em">×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord"><spanclass="mord">×</span></span></span></span></span></span>CH<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo lspace="0em" rspace="0em">×</mo></mrow><annotation encoding="application/x-tex">{\times}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord"><span class="mord">×</span></span></span></span></span></span>W<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo lspace="0em" rspace="0em">×</mo></mrow><annotation encoding="application/x-tex">{\times}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord"><span class="mord">×</span></span></span></span></span></span>C, we compute query (Q\bm{Q}), key (K\bm{K}), and value (V\bm{V}) projections using three linear layers, yielding Q=WQY\bm{Q}{=}\bm{W}^{Q}\bm{Y}, K=WKY\bm{K}{=}\bm{W}^{K}\bm{Y}, and V=WVY\bm{V}{=}\bm{W}^{V}\bm{Y} , with dimensions HW<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><molspace="0em"rspace="0em">×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord"><spanclass="mord">×</span></span></span></span></span></span>CHW<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo lspace="0em" rspace="0em">×</mo></mrow><annotation encoding="application/x-tex">{\times}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord"><span class="mord">×</span></span></span></span></span></span>C, where WQ\bm{W}^{Q},WK\bm{W}^{K}, and WV\bm{W}^{V} are the projection weights for Q\bm{Q}, K\bm{K}, and V\bm{V} respectively. Then, L2 norm is applied to Q\bm{Q} and K\bm{K} before computing the cross-covariance attention as it stabilizes the training. Instead of applying the dot-product between Q\bm{Q} and KT\bm{K}^{T} along the spatial dimension i.e., (HWHW ×{\times} CC) ⋅\cdot (CC ×{\times} HWHW), we apply the dot-product across the channel dimensions between QT\bm{Q}^{T} and K\bm{K} i.e., (C<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><mo>×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord">×</span></span></span></span></span>HWC<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo>×</mo></mrow><annotation encoding="application/x-tex">\times</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord">×</span></span></span></span></span>HW) ⋅\cdot (HW<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><molspace="0em"rspace="0em">×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord"><spanclass="mord">×</span></span></span></span></span></span>CHW<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo lspace="0em" rspace="0em">×</mo></mrow><annotation encoding="application/x-tex">{\times}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord"><span class="mord">×</span></span></span></span></span></span>C), producing C<spanclass="katex−display"><spanclass="katex"><spanclass="katex−mathml"><mathxmlns="http://www.w3.org/1998/Math/MathML"display="block"><semantics><mrow><molspace="0em"rspace="0em">×</mo></mrow><annotationencoding="application/x−tex">×</annotation></semantics></math></span><spanclass="katex−html"aria−hidden="true"><spanclass="base"><spanclass="strut"style="height:0.6667em;vertical−align:−0.0833em;"></span><spanclass="mord"><spanclass="mord">×</span></span></span></span></span></span>CC<span class="katex-display"><span class="katex"><span class="katex-mathml"><math xmlns="http://www.w3.org/1998/Math/MathML" display="block"><semantics><mrow><mo lspace="0em" rspace="0em">×</mo></mrow><annotation encoding="application/x-tex">{\times}</annotation></semantics></math></span><span class="katex-html" aria-hidden="true"><span class="base"><span class="strut" style="height:0.6667em;vertical-align:-0.0833em;"></span><span class="mord"><span class="mord">×</span></span></span></span></span></span>C softmax scaled attention score matrix. To get the final attention maps, we multiply the scores by V\bm{V} and sum them up. The transposed attention operation can be expressed as follows:

where X\bm{X} is the input and X^\hat{\bm{X}} is the output feature tensor. After that, two 1×11{\times}1 pointwise convolution layers, LN and GELU activation are used to generate non-linear features. Table 1 shows the sequence of Conv. and STDA encoders with the corresponding input size at each layer with more design details about extra-extra small, extra-small and small models.

Experiments

In this section, we evaluate our EdgeNeXt model on ImageNet-1K classification, COCO object detection, and Pascal VOC segmentation benchmarks.

We use ImageNet-1K dataset in all classification experiments. The dataset provides approximately 1.28M training and 50K validation images for 1000 categories. Following the literature , we report top-1 accuracy on the validation set for all experiments. For object detection, we use COCO dataset which provides approximately 118k training and 5k validation images respectively. For segmentation, we use Pascal VOC 2012 dataset which provides almost 10k images with semantic segmentation masks. Following the standard practice as in , we use extra data and annotations from and as well.

2 Implementation Details

We train our EdgeNeXt models at an input resolution of 256×\times256 with an effective batch size of 4096. All the experiments are run for 300 epochs with AdamW optimizer, and with a learning rate and weight decay of 6e-3 and 0.05 respectively. We use cosine learning rate schedule with linear warmup for 20 epochs. The data augmentations used during training are Random Resized Crop (RRC), Horizontal Flip, and RandAugment , where RandAugment is only used for the EdgeNeXt-S model. We also use multi-scale sampler during training. Further stochastic depth with a rate of 0.1 is used for EdgeNeXt-S model only. We use EMA with a momentum of 0.9995 during training. For inference, the images are resized to 292×{\times}292 followed by a center crop at 256×{\times}256 resolution. We also train and report the accuracy of our EdgeNeXt-S model at 224×{\times}224 resolution for a fair comparison with previous methods. The classification experiments are run on eight A100 GPUs with an average training time of almost 30 hours for the EdgeNeXt-S model.

For detection and segmentation tasks, we finetune EdgeNeXt following similar settings as in and report mean average precision (mAP) at IOU of 0.50-0.95 and mean intersection over union (mIOU) respectively. The experiments are run on four A100 GPUs with an average training time of ∼\sim36 and ∼\sim7 hours for detection and segmentation respectively.

We also report the latency of our models on NVIDIA Jetson Nanohttps://developer.nvidia.com/embedded/jetson-nano-developer-kit and NVIDIA A100 40GB GPU. For Jetson Nano, we convert all the models to TensorRThttps://github.com/NVIDIA/TensorRT engines and perform inference in FP16 mode using a batch size of 1. For A100, similar to , we use PyTorch v1.8.1 with a batch size of 256 to measure the latency.

3 Image Classification

Table 2 compares our proposed EdgeNeXt model with previous state-of-the-art fully convolutional (ConvNets), transformer-based (ViTs) and hybrid models. Overall, our model demonstrates better accuracy versus compute (parameters and MAdds) trade-off compared to all three categories of methods (see Fig. 1).

Comparison with ConvNets. EdgeNeXt surpasses ligh-weight ConvNets by a formidable margin in terms of top-1 accuracy with similar parameters (Table 2). Normally, ConvNets have less MAdds compared to transformer and hybrid models because of no attention computation, however, they lack the global receptive field. For instance, EdgeNeXt-S has higher MAdds compared to MobileNetV2 , but it obtains 4.1% gain in top-1 accuracy with less number of parameters. Also, our EdgeNeXt-S outperforms ShuffleNetV2 and MobileNetV3 by 4.3% and 3.6% respectively, with comparable number of parameters.

Comparison with ViTs. Our EdgeNeXt outperforms recent ViT variants on ImageNet1K dataset with fewer parameters and MAdds. For example, EdgeNeXt-S obtains 78.8% top-1 accuracy, surpassing T2T-ViT and DeiT-T by 2.3% and 6.6% absolute margins respectively.

Comparison with Hybrid Models. The proposed EdgeNeXt outperforms MobileFormer , ViT-C , CoaT-Lite-T with less number of parameters and fewer MAdds (Table 2). For a fair comparison with MobileViT , we train our model at an input resolution of 256×\times256 and show consistent gains for different models sizes (i.e., S, XS, and XXS) with fewer MAdds and faster inference on the edge devices (Table 3). For instance, our EdgeNeXt-XXS model achieves 71.2% top-1 accuracy with only 1.3M parameters, surpassing the corresponding MobileViT version by 2.2%. Finally, our EdgeNeXt-S model attains 79.4% accuracy on ImageNet with only 5.6M parameters, a margin of 1.0% as compared to the corresponding MobileViT-S model. This demonstrates the effectiveness and the generalization of our design.

Further, we also train our EdgeNeXt-S model using knowledge distillation following and achieves 81.1% top-1 ImageNet accuracy.

4 ImageNet-21K Pretraining

To further explore the capacity of EdgeNeXt, we designed EdgeNeXt-B model with 18.5M parameters and 3.8MAdds and pretrain it on a subset of ImageNet-21K dataset followed by finetuning on standard ImageNet-1K dataset. ImageNet-21K (winter’21 release) contains around 13M images and 19K classes. We follow to preprocess the pretraining data by removing classes with fewer examples and split it into training and validation sets containing around 11M and 522K images respectively over 10,450 classes. We refer this dataset as ImageNet-21K-P. We strictly follow the training recipes of for ImageNet-21K-P pretaining. Further, we initialize the ImageNet-21K-P training with ImageNet-1K pretrained model for faster convergence. Finally, we finetune ImageNet-21K model on ImageNet-1K for 30 epochs with a learning rate of 7.5e-5 and an effective batch size of 512. The results are summarized in Table 4.

5 Inference on Edge Devices

We compute the inference time of our EdgeNeXt models on the NVIDIA Jetson Nano edge device and compare it with the state-of-the-art MobileViT model (Table 3). All the models are converted to TensorRT engines and inference is performed in FP16 mode. Our model attains low latency on the edge device with similar parameters, fewer MAdds, and higher top-1 accuracy. Table 3 also lists the inference time on A100 GPU for both MobileViT and EdgeNeXt models. It can be observed that our EdgeNeXt-XXS model is ∼\sim34% faster than the MobileViT-XSS model on A100 as compared to only ∼\sim8% faster on Jetson Nano, indicating that EdgeNeXt better utilizes the advanced hardware as compared to MobileViT.

6 Object Detection

We use EdgeNeXt as a backbone in SSDLite and finetune the model on COCO 2017 dataset at an input resolution of 320×{\times}320. The difference between SSD and SSDLite is that the standard convolutions are replaced with separable convolutions in the SSD head. The results are reported in Table 5. EdgeNeXt consistently outperforms MobileNet backbones and gives competitive performance compared to MobileVit backbone. With less number of MAdds and a comparable number of parameters, EdgeNeXt achieves the highest 27.9 box AP, ∼\sim38% fewer MAdds than MobileViT.

7 Semantic Segmentation

We use EdgeNeXt as backbone in DeepLabv3 and finetune the model on Pascal VOC dataset at an input resolution of 512×\times512. DeepLabv3 uses dilated convolution in cascade design along with spatial pyramid pooling to encode multi-scale features which are useful in encoding objects at multiple scales. Our model obtains 80.2 mIOU on the validation dataset, providing a 1.1 points gain over MobileViT with ∼\sim36% fewer MAdds.

Ablations

In this section, we ablate different design choices in our proposed EdgeNeXt model.

SDTA encoder and adaptive kernel sizes. Table 7 illustrates the importance of SDTA encoders and adaptive kernel sizes in our proposed architecture. Replacing SDTA encoders with convolution encoders degrades the accuracy by 1.1%, indicating the usefulness of SDTA encoders in our design. When we fix kernel size to 7 in all four stages of the network, it further reduces the accuracy by 0.4%. Overall, our proposed design provides an optimal speed-accuracy trade-off.

We also ablate the contributions of SDTA components (e.g., adaptive branching and positional encoding) in Table 7. Removing adaptive branching and positional encoding slightly decreases the accuracy.

Hybrid design. Table 8 ablates the different hybrid design choices for our EdgeNeXt model. Motivated from MetaFormer , we replace all convolutional modules in the last two stages with SDTA encoders. The results show superior performance when all blocks in the last two stages are SDTA blocks, but it increases the latency (row-2 vs 3). Our hybrid design where we propose to use an SDTA module as the last block in the last three stages provides an optimal speed-accuracy trade-off.

Table 9 provides an ablation of the importance of using SDTA encoders at different stages of the network. It is noticable that progressively adding an SDTA encoder as the last block of the last three stages improves the accuracy with some loss in inference latency. However, in row 4, we obtain the best trade-off between accuracy and speed where the SDTA encoder is added as the last block in the last three stages of the network. Further, we notice that adding a global SDTA encoder to the first stage of the network is not helpful where the features are not much mature.

We also provide an ablation on using the SDTA module at the start of each stage versus at the end. Table 10 shows that using the global SDTA encoder at the end of each stage is more beneficial. This observation is consistent with the recent work .

Activation and normalization. EdgeNeXt uses GELU activation and layer normalization throughout the network. We found that the current PyTorch implementations of GELU and layer normalization are not optimal for high speed inference. To this end, we replace GELU with Hard-Swish and layer-norm with batch-norm and retrain our models. Fig. 3 indicates that it reduces the accuracy slightly, however, reduces the latency by a large margin.

Qualitative Results

Figs. 4 and 5 shows the qualitative results of EdgeNeXt detection and segmentation models respectively. Our model can detect and segment objects in various views.

Conclusion

The success of the transformer models comes with a higher computational overhead compared to CNNs. Self-attention operation is the major contributor to this overhead, which makes vision transformers slow on the edge devices compared to CNN-based mobile architectures. In this paper, we introduce a hybrid design consisting of convolution and efficient self-attention based encoders to jointly model local and global information effectively, while being efficient in terms of both parameters and MAdds on vision tasks with superior performance compared to state-of-the-art methods. Our experimental results show promising performance for different variants of EdgeNeXt, which demonstrates the effectiveness and the generalization ability of the proposed model.

References