Conditional DETR for Fast Training Convergence

Depu Meng, Xiaokang Chen, Zejia Fan, Gang Zeng, Houqiang Li, Yuhui Yuan, Lei Sun, Jingdong Wang

Introduction

The DEtection TRansformer (DETR) method applies the transformer encoder and decoder architecture to object detection and achieves good performance. It effectively eliminates the need for many hand-crafted components, including non-maximum suppression and anchor generation.

The DETR approach suffers from slow convergence on training, and needs 500500 training epochs to get good performance. The very recent work, deformable DETR , handles this issue by replacing the global dense attention (self-attention and cross-attention) with deformable attention that attends to a small set of key sampling points and using the high-resolution and multi-scale encoder. Instead, we still use the global dense attention and propose an improved decoder cross-attention mechanism for accelerating the training process.

Our approach is motivated by high dependence on content embeddings and minor contributions made by the spatial embeddings in cross-attention. The empirical results in DETR show that if removing the positional embeddings in keys and the object queries from the second decoder layer and only using the content embeddings in keys and queries, the detection AP drops slightly The minor AP drop 1.41.4 is reported on R5050 with 300300 epochs in Table 3 from . We empirically got the consistent observation: the AP drops to 34.034.0 from 34.934.9 for 5050 training epochs. .

Figure 1 (the second row) shows that the spatial attention weight maps from the cross-attention in DETR trained with 5050 epochs. One can see that two among the four maps do not correctly highlight the bands for the corresponding extremities, thus weak at shrinking the spatial range for the content queries to precisely localize the extremities. The reasons are that (i) the spatial queries, i.e., object queries, only give the general attention weight map without exploiting the specific image information; and that (ii) due to short training the content queries are not strong enough to match the spatial keys well as they are also used to match the content keys. This increases the dependence on high-quality content embeddings, thus increasing the training difficulty.

We present a conditional DETR approach, which learns a conditional spatial embedding for each query from the corresponding previous decoder output embedding, to form a so-called conditional spatial query for decoder multi-head cross-attention. The conditional spatial query is predicted by mapping the information for regressing the object box to the embedding space, the same to the space that the 22D coordinates of the keys are also mapped to.

We empirically observe that using the spatial queries and keys, each cross-attention head spatially attends to a band containing the object extremity or a region inside the object box (Figure 1, the first row). This shrinks the spatial range for the content queries to localize the effective regions for class and box prediction. As a result, the dependence on the content embeddings is relaxed and the training is easier. The experiments show that conditional DETR converges 6.7×6.7\times faster for the backbones R5050 and R101101 and 10×10\times faster for stronger backbones DC55-R5050 and DC55-R101101. Figure 2 gives the convergence curves for conditional DETR and the original DETR .

Related Work

Anchor-based and anchor-free detection. Most existing object detection approaches make predictions from initial guesses that are carefully designed. There are two main initial guesses: anchor boxes or object centers. The anchor box-based methods inherit the ideas from the proposal-based method, Fast R-CNN. Example methods include Faster R-CNN , SSD , YOLOv2 , YOLOv3 , YOLOv4 , RetinaNet , Cascade R-CNN , Libra R-CNN , TSD and so on.

The anchor-free detectors predict the boxes at points near the object centers. Typical methods include YOLOv1 , CornerNet , ExtremeNet , CenterNet , FCOS and others .

DETR and its variants. DETR successfully applies transformers to object detection, effectively removing the need for many hand-designed components like non-maximum suppression or initial guess generation. The high computation complexity issue, caused by the global encoder self-attention, is handled in adaptive clustering transformer and by sparse attentions in deformable DETR .

The other critical issue, slow training convergence, has been attracting a lot of recent research attention. The TSP (transformer-based set prediction) approach eliminates the cross-attention modules and combines the FCOS and R-CNN-like detection heads. Deformable DETR adopts deformable attention, which attends to sparse positions learned from the content embedding, to replace decoder cross-attention.

The spatially modulated co-attention (SMCA) approach , which is concurrent to our approach, is very close to our approach. It modulates the DETR multi-head global cross-attentions with Gaussian maps around a few (shifted) centers that are learned from the decoder embeddings, to focus more on a few regions inside the estimated box. In contrast, the proposed conditional DETR approach learns the conditional spatial queries from the decoder content embeddings, and predicts the spatial attention weight maps without human-crafting the attention attenuation, which highlight four extremities for box regression, and distinct regions inside the object for classification.

Conditional and dynamic convolution. The proposed conditional spatial query scheme is related to conditional convolutional kernel generation. Dynamic filter network learns the convolutional kernels from the input, which is applied to instance segmentation in CondInst and SOLOv2 for learning instance-dependent convolutional kernels. CondConv and dynamic convolution mix convolutional kernels with the weights learned from the input. SENet , GENet abd Lite-HRNet learn from the input the channel-wise weights.

These methods learn from the input the convolutional kernel weights and then apply the convolutions to the input. In contrast, the linear projection in our approach is learned from the decoder embeddings for representing the displacement and scaling information.

Transformers. The transformer relies on the attention mechanism, self-attention and cross-attention, to draw global dependencies between the input and the output. There are several works closely related to our approach. Gaussian transformer and T-GSA (Transformer with Gaussian-weighted self-attention) , followed by SMCA , attenuate the attention weights according to the distance between target and context symbols with learned or human-crafted Gaussian variance. Similar to ours, TUPE computes the attention weight also from the spatial attention weight and the content attention weight. Instead, our approach mainly focuses on the attention attenuation mechanism in a learnable form other than a Gaussian function, and potentially benefits speech enhancement and natural language inference .

Conditional DETR

Pipeline. The proposed approach follows detection transformer (DETR), an end-to-end object detector, and predicts all the objects at once without the need for NMS or anchor generation. The architecture consists of a CNN backbone, a transformer encoder, a transformer decoder, and object class and box position predictors. The transformer encoder aims to improve the content embeddings output from the CNN backbone. It is a stack of multiple encoder layers, where each layer mainly consists of a self-attention layer and a feed-forward layer.

The transformer decoder is a stack of decoder layers. Each decoder layer, illustrated in Figure 3, is composed of three main layers: (1) a self-attention layer for removing duplication prediction, which performs interactions between the embeddings, outputted from the previous decoder layer and used for class and box prediction, (2) a cross-attention layer, which aggregates the embeddings output from the encoder to refine the decoder embeddings for improving class and box prediction, and (3) a feed-forward layer.

Box regression. A candidate box is predicted from each decoder embedding as follows,

Here, f\mathbf{f} is the decoder embedding. b\mathbf{b} is a four-dimensional vector [bcx bcy bw bh]⊤[b_{cx}~{}b_{cy}~{}b_{w}~{}b_{h}]^{\top}, consisting of the box center, the box width and the box height. sigmoid⁡()\operatorname{sigmoid}() is used to normalize the prediction b\mathbf{b} to the range $..\operatorname{FFN}()aimstopredicttheunnormalizedbox.aims to predictthe unnormalized box.{\mathbf{s}}istheunnormalizedis the unnormalized2Dcoordinateofthereferencepoint,andisD coordinate of the reference point, and is(0,0)intheoriginalDETR.Inourapproach,weconsidertwochoices:learnthereferencepointin the original DETR. In our approach, we consider two choices: learn the reference point\mathbf{s}$ as a parameter for each candidate box prediction, or generate it from the corresponding object query.

Category prediction. The classification score for each candidate box is also predicted from the decoder embedding through an FNN, e=FFN⁡(f)\mathbf{e}=\operatorname{FFN}(\mathbf{f}).

Main work. The cross-attention mechanism aims to localize the distinct regions, four extremities for box detection and regions inside the box for object classification, and aggregates the corresponding embeddings. We propose a conditional cross-attention mechanism with introducing conditional spatial queries for improving the localization capability and accelerating the training process.

2 DETR Decoder Cross-Attention

The DETR decoder cross-attention mechanism takes three inputs: queries, keys and values. Each key is formed by adding a content key ck\mathbf{c}_{k} (the content embedding output from the encoder) and a spatial key pk\mathbf{p}_{k} (the positional embedding of the corresponding normalized 22D coordinate). The value is formed from the content embedding, same with the content key, output from the encoder.

In the original DETR approach, each query is formed by adding a content query cq\mathbf{c}_{q} (the embedding output from the decoder self-attention), and a spatial query pq\mathbf{p}_{q} (i.e., the object query oq\mathbf{o}_{q}). In our implementation, there are N=300N=300 object queries, and accordingly there are NN queriesFor description simplicity and clearness, we drop the query, key, and value indices., each query outputting a candidate detect result in one decoder layer.

The attention weight is based on the dot-product between the query and the key, used for attention weight computation,

3 Conditional Cross-Attention

The proposed conditional cross-attention mechanism forms the query by concatenating the content query cq\mathbf{c}_{q}, outputting from decoder self-attention, and the spatial query pq\mathbf{p}_{q}. Accordingly, the key is formed as the concatenation of the content key ck\mathbf{c}_{k} and the spatial key pk\mathbf{p}_{k}.

The cross-attention weights consist of two components, content attention weight and spatial attention weight. The two weights are from two dot-products, content and spatial dot-products,

Different from the original DETR cross-attention, our mechanism separates the roles of content and spatial queries so that spatial and content queries focus on the spatial and content attention weights, respectively.

An additional important task is to compute the spatial query pq\mathbf{p}_{q} from the embedding f\mathbf{f} of the previous decoder layer. We first identify that the spatial information of the distinct regions are determined by the two factors together, decoder embedding and reference point. We then show how to map them to the embedding space, forming the query pq\mathbf{p}_{q}, so that the spatial query lies in the same space the 22D coordinates of the keys are mapped to.

The decoder embedding contains the displacements of the distinct regions with respect to the reference point. The box prediction process in Equation 1 consists of two steps: (1) predicting the box with respect to the reference point in the unnormalized space, and (2) normalizing the predicted box to the range $TheoriginThe origin(0,0)intheunnormalizedspacefortheoriginalDETRmethodismappedtoin the unnormalized space for the original DETR method is mapped to(0.5,0.5)(thecenterintheimagespace)inthenormalizedspacethroughthe(the center in the image space) in the normalized space through the\operatorname{sigmoid}$ function..

Step (1) means that the decoder embedding f\mathbf{f} contains the displacements of the four extremities (forming the box) with respect to the reference point s\mathbf{s} in the unnormalized space. This implies that both the embedding f\mathbf{f} and the reference point s\mathbf{s} are necessary to determine the spatial information of the distinct regions, the four extremities as well as the region for predicting the classification score.

Conditional spatial query prediction. We predict the conditional spatial query from the embedding f\mathbf{f} and the reference point s\mathbf{s},

so that it is aligned with the positional space which the normalized 22D coordinates of the keys are mapped to. The process is illustrated in the gray-shaded box area of Figure 3.

We normalize the reference point s\mathbf{s} and then map it to a 256256-dimensional sinusoidal positional embedding in the same way as the positional embedding for keys:

We then map the displacement information contained in the decoder embedding f\mathbf{f} to a linear projection in the same space through an FFN consisting of learnable linear projection + ReLU + learnable linear projection: T=FFN⁡(f)\mathbf{T}=\operatorname{FFN}(\mathbf{f}).

The conditional spatial query is computed by transforming the reference point in the embedding space: pq=Tps\mathbf{p}_{q}=\mathbf{T}\mathbf{p}_{s}. We choose the simple and computationally-efficient projection matrix, a diagonal matrix. The 256256 diagonal elements are denoted as a vector \uplambdaq\boldsymbol{\uplambda}_{q}. The conditional spatial query is computed by the element-wise multiplication:

Multi-head cross-attention. Following DETR , we adopt the standard multi-head cross-attention mechanism. Object detection usually needs to implicitly or explicitly localize the four object extremities for accurate box regression and localize the object region for accurate object classification. The multi-head mechanism is beneficial to disentangle the localization tasks.

We perform multi-head parallel attentions by projecting the queries, the keys, and the values M=8M=8 times with learned linear projections to low dimensions. The spatial and content queries (keys) are separately projected to each head with different linear projections. The projections for values are the same as the original DETR and are only for the contents.

4 Visualization and Analysis

Visualization. Figure 4 visualizes the attention weight maps for each head: the spatial attention weight maps, the content attention weight maps, and the combined attention weight maps. The maps are soft-max normalized over the spatial dot-products pq⊤pk\mathbf{p}_{q}^{\top}\mathbf{p}_{k}, the content dot-products cq⊤ck\mathbf{c}_{q}^{\top}\mathbf{c}_{k}, and the combined dot-products cq⊤ck+pq⊤pk\mathbf{c}_{q}^{\top}\mathbf{c}_{k}+\mathbf{p}_{q}^{\top}\mathbf{p}_{k}. We show 55 out of the 88 maps, and other three are the duplicates, corresponding to bottom and top extremities, and a small region inside the object boxThe duplicates might be different for models trained several times, but the detection performance is almost the same..

We can see that the spatial attention weight map at each head is able to localize a distinct region, a region containing one extremity or a region inside the object box. It is interesting that each spatial attention weight map corresponding to an extremity highlights a spatial band that overlaps with the corresponding edge of the object box. The other spatial attention map for the region inside the object box merely highlights a small region whose representations might already encode enough information for object classification.

The content attention weight maps of the four heads corresponding to the four extremities highlight scattered regions in addition to the extremities. The combination of the spatial and content maps filters out other highlights and keeps extremity highlights for accurate box regression.

Comparison to DETR. Figure 1 shows the spatial attention weight maps of our conditional DETR (the first row) and the original DETR trained with 5050 epochs (the second row). The maps of our approach are computed by soft-max normalizing the dot-products between spatial keys and queries, pq⊤pk\mathbf{p}_{q}^{\top}\mathbf{p_{k}}. The maps for DETR are computed by soft-max normalizing the dot-products with the spatial keys, (oq+cq)⊤pk(\mathbf{o}_{q}+\mathbf{c}_{q})^{\top}\mathbf{p}_{k}.

It can be seen that our spatial attention weight maps accurately localize the distinct regions, four extremities. In contrast, the maps from the original DETR with 5050 epochs can not accurately localize two extremities, and 500500 training epochs (the third row) make the content queries stronger, leading to accurate localization. This implies that it is really hard to learn the content query cq\mathbf{c}_{q} to serve as two rolesStrictly speaking, the embedding output from decoder self-attention for more training epochs contains both spatial and content information. For discussion convenience, we still call it content query.: match the content key and the spatial key simultaneously, and thus more training epochs are needed.

Analysis. The spatial attention weight maps shown in Figure 4 imply that the conditional spatial query, used to form the spatial query, have at least two effects. (i) Translate the highlight positions to the four extremities and the position inside the object box: interestingly the highlighted positions are spatially similarly distributed in the object box. (ii) Scale the spatial spread for the extremity highlights: large spread for large objects and small spread for small objects.

The two effects are realized in the spatial embedding space through applying the transformation T\mathbf{T} over ps\mathbf{p}_{s} (further disentangled through image-independent linear projections contained in cross-attention and distributed to each head). This indicates that the transformation T\mathbf{T} not only contains the displacements as discussed before, but also the object scale.

5 Implementation Details

Architecture. Our architecture is almost the same with the DETR architecture and contains the CNN backbone, transformer encoder, transformer decoder, prediction feed-forward networks (FFNs) following each decoder layer (the last decoder layer and the 55 internal decoder layers) with parameters shared among the 66 prediction FFNs. The hyper-parameters are the same as DETR.

The main architecture difference is that we introduce the conditional spatial embeddings as the spatial queries for conditional multi-head cross-attention and that the spatial query (key) and the content query (key) are combined through concatenation other than addition. In the first cross-attention layer there are no decoder content embeddings, we make simple changes based on the DETR implementation : concatenate the positional embedding predicted from the object query (the positional embedding) into the original query (key).

Reference points. In the original DETR approach, s=[0 0]⊤{\mathbf{s}}=[0~{}0]^{\top} is the same for all the decoder embeddings. We study two ways forming the reference points: regard the unnormalized 22D coordinates as learnable parameters, and the unnormalized 22D coordinate predicted from the object query oq\mathbf{o}_{q}. In the latter way that is similar to deformable DETR , the prediction unit is an FFN and consists of learnable linear projection + ReLU + learnable linear projection: s=FFN⁡(oq)\mathbf{s}=\operatorname{FFN}(\mathbf{o}_{q}). When used for forming the conditional spatial query, the 22D coordinates are normalized by the sigmoid function.

Loss function. We follow DETR to find an optimal bipartite matching between the predicted and ground-truth objects using the Hungarian algorithm, and then form the loss function for computing and back-propagate the gradients. We use the same way with deformable DETR to formulate the loss: the same matching cost function, the same loss function with 300300 object queries, and the same trade-off parameters; The classification loss function is focal loss , and the box regression loss (including L1 and GIoU loss) is the same as DETR .

Experiments

Dataset. We perform the experiments on the COCO 20172017 detection dataset. The dataset contains about 118118K training images and 55K validation (val) images.

Training. We follow the DETR training protocol . The backbone is the ImageNet-pretrained model from TORCHVISION with batchnorm layers fixed, and the transformer parameters are initialized using the Xavier initialization scheme . The weight decay is set to be 10−410^{-4}. The AdamW optimizer is used. The learning rates for the backbone and the transformer are initially set to be 10−510^{-5} and 10−410^{-4}, respectively. The dropout rate in transformer is 0.10.1. The learning rate is dropped by a factor of 1010 after 4040 epochs for 5050 training epochs, after 6060 epochs for 7575 training epochs, and after 8080 epochs for 108108 training epochs.

We use the augmentation scheme same as DETR : resize the input image such that the short side is at least 480480 and at most 800800 pixels and the long size is at most 13331333 pixels; randomly crop the image such that a training image is cropped with probability 0.50.5 to a random rectangular patch.

Evaluation. We use the standard COCO evaluation. We report the average precision (AP), and the AP scores at 0.500.50, 0.750.75 and for the small, medium, and large objects.

2 Results

Comparison to DETR. We compare the proposed conditional DETR to the original DETR . We follow and report the results over four backbones: ResNet-5050 , ResNet-101101, and their 16×16\times-resolution extensions DC55-ResNet-5050 and DC55-ResNet-101101.

The corresponding DETR models are named as DETR-R5050, DETR-R101101, DETR-DC5-R5050, and DETR-DC5-R101101, respectively. Our models are named as conditional DETR-R5050, conditional DETR-R101101, conditional DETR-DC5-R5050, and conditional DETR-DC5-R101101, respectively.

Table 1 presents the results from DETR and conditional DETR. DETR with 5050 training epochs performs much worse than 500500 training epochs. Conditional DETR with 5050 training epochs for R5050 and R101101 as the backbones performs slightly worse than DETR with 500500 training epochs. Conditional DETR with 5050 training epochs for DC55-R5050 and DC55-R101101 performs similarly as DETR with 500500 training epochs. Conditional DETR for the four backbones with 75/10875/108 training epochs performs better than DETR with 500500 training epochs. In summary, conditional DETR for high-resolution backbones DC55-R5050 and DC55-R101101 is 10×10\times faster than the original DETR, and for low-resolution backbones R5050 and R101101 6.67×6.67\times faster. In other words, conditional DETR performs better for stronger backbones with better performance.

In addition, we report the results of single-scale DETR extensions: deformable DETR-SS and UP-DETR in Table 1. Our results over R5050 and DC55-R5050 are better than deformable DETR-SS: 40.940.9 vs. 39.439.4 and 43.843.8 vs. 41.541.5. The comparison might not be fully fair as for example parameter and computation complexities are different, but it implies that the conditional cross-attention mechanism is beneficial. Compared to UP-DETR-R5050, our results with fewer training epochs are obviously better.

Comparison to multi-scale and higher-resolution DETR variants. We focus on accelerating the DETR training, without addressing the issue of high computational complexity in the encoder. We do not expect that our approach achieves on par with DETR variants w/ multi-scale attention and 8×8\times-resolution encoders, e.g., TSP-FCOS and TSP-RCNN and deformable DETR , which are able to reduce the encoder computational complexity and improve the performance due to multi-scale and higher-resolution.

The comparisons in Table 2 surprisingly show that our approach on DC55-R5050 (16×16\times) performs same as deformable DETR-R5050 (multi-scale, 8×8\times). Considering that the AP of the single-scale deformable DETR-DC55-R5050-SS is 41.541.5 (lower than ours 43.843.8) (Table 1), one can see that deformable DETR benefits a lot from the multi-scale and higher-resolution encoder that potentially benefit our approach, which is currently not our focus and left as our future work.

The performance of our approach is also on par with TSP-FCOS and TSP-RCNN. The two methods contain a transformer encoder over a small number of selected positions/regions (feature of interest in TSP-FCOS and region proposals in TSP-RCNN) without using the transformer decoder, are extensions of FCOS and Faster RCNN . It should be noted that position/region selection removes unnecessary computation in self-attention and reduces computation complexity dramatically.

3 Ablations

Reference points. We compare three ways of forming reference points s\mathbf{s}: (i) s=(0,0)\mathbf{s}=(0,0), same to the original DETR, (ii) learn s\mathbf{s} as model parameters and each prediction is associated with different reference points, and (iii) predict each reference point s\mathbf{s} from the corresponding object query. We conducted the experiments with ResNet-5050 as the backbone. The AP scores are 36.836.8, 40.740.7, and 40.940.9, suggesting that (ii) and (iii) perform on par and better than (i).

The effect of the way forming the conditional spatial query. We empirically study how the transformation \uplambdaq\boldsymbol{\uplambda}_{q} and the positional embedding ps\mathbf{p}_{s} of the reference point, used to form the conditional spatial query pq=\uplambdaq⊙ps\mathbf{p}_{q}=\boldsymbol{\uplambda}_{q}\odot\mathbf{p}_{s}, make contributions to the detection performance.

We report the results of our conditional DETR, and the other ways forming the spatial query with: (i) CSQ-P - only the positional embedding ps\mathbf{p}_{s}, (ii) CSQ-T - only the transformation \uplambdaq\boldsymbol{\uplambda}_{q}, (iii) CSQ-C - the decoder content embedding f\mathbf{f}, and (iv) CSQ-I - the element-wise product of the transformation predicted from the decoder self-attention output cq\mathbf{c}_{q} and the positional embedding ps\mathbf{p}_{s}. The studies in Table 3 imply that our proposed way (CSQ) performs overall the best, validating our analysis about the transformation predicted from the decoder embedding and the positional embedding of the reference point in Section 3.3.

Focal loss and offset regression with respect to learned reference point. Our approach follows deformable DETR : use the focal loss with 300300 object queries to form the classification loss and predict the box center by regressing the offset with respect to the reference point. We report how the two schemes affect the DETR performance in Table 4. One can see that separately using the focal loss or center offset regression without learning referecence points leads to a slight AP gain and combining them together leads to a larger AP gain. Conditional cross-attention in our approach built on the basis of focal loss and offset regression brings a major gain 4.04.0.

The effect of linear projections T\mathbf{T} forming the transformation. Predicting the conditional spatial query needs to learn the linear projection T\mathbf{T} from the decoder embedding (see Equation 6). We empirically study how the linear projection forms affect the performance. The linear projection forms include: an identity matrix that means not to learn the linear projection, a single scalar, a block diagonal matrix meaning that each head has a learned 32×3232\times 32 linear projection matrix, a full matrix without constraints, and a diagonal matrix. Figure 5 presents the results. It is interesting that a single-scalar helps improve the performance, maybe due to narrowing down the spatial range to the object area. Other three forms, block diagonal, full, and diagonal (ours), perform on par.

Conclusion

We present a simple conditional cross-attention mechanism. The key is to learn a spatial query from the corresponding reference point and decoder embedding. The spatial query contains the spatial information mined for the class and box prediction in the previous decoder layer, and leads to spatial attention weight maps highlighting the bands containing extremities and small regions inside the object box. This shrinks the spatial range for the content query to localize the distinct regions, thus relaxing the dependence on the content query and reducing the training difficulty. In the future, we will study the proposed conditional cross-attention mechanism for human pose estimation and line segment detection .

Acknowledgments. We thank the anonymous reviewers for their insightful comments and suggestions on our manuscript.

References