Stratified Transformer for 3D Point Cloud Segmentation

Xin Lai, Jianhui Liu, Li Jiang, Liwei Wang, Hengshuang Zhao, Shu Liu, Xiaojuan Qi, Jiaya Jia

Introduction

Nowadays 3D point clouds can be conveniently collected. They have demonstrated great potential in various applications, such as autonomous driving, robotics and augmented reality. Unlike regular pixels in 2D images, 3D points are arranged irregularly, hampering direct adoption of well-studied 2D networks to process 3D data. Therefore, it is imperative to explore advanced methods that are tailored for 3D point cloud data.

Abundant methods have explored 3D point cloud segmentation and obtained decent performance. Most of them focus on aggregating local features, but fail to explicitly model long-range dependencies, which has been demonstrated to be crucial in capturing contexts from a long distance . Along another line of research, Transformer can naturally harvest long-range information via the self-attention mechanism. However, only limited attempts have been made to apply Transformer to 3D point clouds. Point Transformer proposes “vector self-attention” and “subtraction relation” to aggregate local features, but it is still difficult to directly capture long-range contexts. Voxel Transformer is tailored for object detection and performs self-attention over the voxels, but it loses accurate position due to voxelization.

Differently, we develop an efficient segmentation network to capture long-range contexts using the standard multi-head self-attention , while keeping position information intact. To this end, we propose a simple and powerful framework, namely, Stratified Transformer.

Specifically, we first partition the 3D space into non-overlapping cubic windows, inspired by Swin Transformer . However, in Swin Transformer, different windows work independently, and each query token only chooses the tokens within its window as keys, thus attending to a limited local region. Instead, we propose a stratified strategy for sampling keys. Rather than only selecting nearby points in the same window as keys, we also sparsely sample distant points. In this way, for each query point, both denser nearby points and sparser distant points are sampled to form the keys all together, achieving a significantly enlarged effective receptive field while incurring negligible extra computations. For instance, we visualize the Effective Receptive Field (ERF) in Fig. 1 to show the importance of modeling long-range contexts. In the middle of the figure, due to incapability to model the direct long-range dependency, the desk merely attends to the local region, leading to false predictions. Contrarily, with our proposed stratified strategy, the desk is able to aggregate contexts from distant objects, such as the bed or curtain, which helps to correct the prediction.

Moreover, it is notable that irregular point arrangements pose significant challenges in designing 3D Transformer. In 2D images, patch-wise tokens can be easily formed with spatially regular pixels. But 3D points are completely different. In our framework, each point is deemed as a token and we perform point embedding for each point to aggregate local information in the first layer, which is beneficial for faster convergence and stronger performance. Furthermore, we adopt effective relative position encoding to capture richer position information. It can generate the positional bias dynamically with contexts, through the interaction with the semantic features. Also, considering that 3D point numbers in different windows vary a lot and cause unnecessary memory occupation for windows with a small number of points, we introduce a memory-efficient implementation to significantly reduce memory consumption.

We propose Stratified Transformer to additionally sample distant points as keys but in a sparser way, enlarging the effective receptive field and building direct long-range dependency while incurring negligible extra computations.

To handle irregular point arrangements, we design first-layer point embedding and effective contextual position encoding, along with a memory-efficient implementation, to build a strong Transformer tailored for 3D point cloud segmentation.

Experiments show our model achieves state-of-the-art results on widely adopted large-scale segmentation datasets, i.e., S3DIS , ScanNetv2 and ShapeNetPart . Extensive ablation studies verify the benefit of each component.

Related Work

Recently, vision Transformer becomes popular in 2D image understanding . ViT treats each patch as a token, and directly uses a Transformer encoder to extract features for image classification. Further, PVT proposes a hierarchical structure to obtain a pyramid of features for semantic segmentation and also presents Spatial Reduction Attention to save memory. Alternatively, Swin Transformer uses a window-based attention, and proposes a shifted window operation in the successive Transformer block. Methods of further propose different designs to incorporate long-range and global dependencies. Transformer is already popular in 2D, but remains under-explored on point clouds. Inspired by Swin Transformer, we adopt hierarchical structure and shifted window operation for 3D point cloud. On top of that, we propose a stratified strategy for sampling keys to harvest long-range contexts, and put forward several essential designs to combat the challenges posed by irregular point arrangements.

Point Cloud Segmentation.

Approaches for point cloud segmentation can be grouped into two categories, i.e., voxel-based and the point-based methods. The voxel-based solutions first divide the 3D space into regular voxels, and then apply sparse convolutions upon them. They yield decent performance, but suffer from inaccurate position information due to voxelization. Point-based methods directly adopt the point features and positions as inputs, thus keeping the position information intact. Following this line of research, different ways for feature aggregation are designed to learn high-level semantic features. PointNet and its variants use max pooling to aggregate features. PointConv and KPConv try to use an MLP or discrete kernel points to mimic a continuous convolution kernel. Point Transformer uses the “vector self-attention” operator to aggregate local features and the “subtraction relation” to generate the attention weights, but it suffers from lack of long-range contexts and insufficient robustness upon various perturbations in testing.

Our work is pointed-based and closely related point transformer yet with a fundamental difference: ours overcomes the limited effective receptive field issue and makes the best of Transformer for modeling long-range contextual dependencies instead of merely local aggregation.

Our Method

The overview of our model is illustrated in Fig. 2. Our framework is point-based, and we use both xyz coordinates and rgb colors as input. The encoder-decoder structure is adopted where the encoder is composed of multiple stages connected by downsample layers. At the beginning of the encoder, the first-layer point embedding module is used for local aggregation. Then, there are several Transformer blocks at each stage. As for the decoder, the encoder features are upsampled to become denser layer by layer in the way similar to U-Net .

2 Transformer Block

The Transformer block is composed of a standard multi-head self-attention module and a feed-forward network (FFN). With tens of thousands of points as inputs, directly applying global self-attention incurs unacceptable O(N2)O(N^{2}) memory consumption, where NN is the input point number.

Note that the above equations only show the calculation in a single window, and different windows work in the same way independently. In this way, the memory complexity is dramatically reduced to O(Nk×k2)=O(N×k)O(\frac{N}{k}\times k^{2})=O(N\times k), where kk is the average number of points scattered in each window.

To facilitate cross-window communication, we also shift the window by half of the window size between two successive Transformer blocks, similar to . The illustration of shifted window is given in the supplementary file.

Stratified Key-sampling Strategy.

Since every query point only attends to the local points in its own window, the vanilla version Transformer block suffers from limited effective receptive field even with shifted window, as shown in Fig. 1. Therefore, it fails to capture long-range contextual dependencies over distant objects, causing false predictions.

A simple solution is to enlarge the size of cubic window. However, the memory would grow as the window size increases. To effectively aggregate long-range contexts at a low cost of memory, we propose a stratified strategy for sampling keys. As shown in Fig. 3, we partition the space into non-overlapping cubic windows with the window size swins_{win}. For each query point qiq_{i} (shown with green star), we find the points Kidense\mathbf{K}_{i}^{dense} in its window, same as the vanilla version. Additionally, we downsample the input points through farthest point sampling (fps) at the scale of ss, and find the points Kisparse\mathbf{K}_{i}^{sparse} with a larger window size swinlarges_{win}^{large}. In the end, both dense and sparse keys form the final keys, i.e., Ki=Kidense∪Kisparse\mathbf{K}_{i}=\mathbf{K}_{i}^{dense}\cup\mathbf{K}_{i}^{sparse}. Note that duplicated key points are only counted once.

The complete structure of Stratified Transformer block is shown in Fig. 2 (b). Following common practice, we use LayerNorm before each self-attention module or feed-forward network. To further complement the information interaction across windows, the original window is shifted by 12swin\frac{1}{2}s_{win} while the large window is shifted by 12swinlarge\frac{1}{2}s_{win}^{large} in the successive Transformer block. This further boosts the performance as listed in Table 7.

Thanks to the stratified strategy for key sampling, the effective receptive field is enlarged remarkably and the query feature is able to effectively aggregate long-range contexts. Compared to the vanilla version, we merely incur the extra computations on the sparse distant keys, which only takes up about 10% of the final keys Ki\mathbf{K}_{i}.

3 First-layer Point Embedding

In the first layer, we build a point embedding module. An intuitive choice is to use a linear layer or MLP to project the input features to a high dimension. However, we empirically observe relatively slow convergence and poor performance by using a linear layer in the first layer, as shown in Fig. 4. We note that the point feature from a linear layer or MLP merely comprises the raw information of its own xyz position and the rgb color, but it lacks local geometric and contextual information. As a result, in the first Transformer block, the attention map could not capture high-level relevance between the queries and keys that only contain raw xyz and rgb information. This negatively affects representation power and generalization ability of the model.

We contrarily propose to aggregate the features of local neighbors for each point in the Point Embedding module. We try a variety of methods for local aggregation, such as max pooling and average pooling, and find KPConv performs the best, as shown in Table 5. Surprisingly, this minor modification to the architecture brings about considerable improvement as suggested in Exp.I and II as well as Exp.V and VI of Table 4. It proves the importance of initial local aggregation in the Transformer-based networks. Note that a single KPConv incurs negligible extra computations (merely 2%2\% FLOPs) compared to the whole network.

4 Contextual Relative Position Encoding

Compared to 2D spatially regular pixels, 3D points are in a more complicated continuous space, posing challenges to exploit the xyz position. claims that position encoding is unnecessary for 3D Transformer-based networks because the xyz coordinates have already been used as the input features. However, although the input of the Transformer block has already contained the xyz position, fine-grained position information may be lost in high-level features when going deeper through the network. To make better use of the position information, we adopt a context-based adaptive relative position encoding scheme inspired by .

where swins_{win} is the window size and squant=2⋅swinLs_{quant}=\frac{2\cdot s_{win}}{L} is the quantization size, and ⌊⋅⌋\lfloor\cdot\rfloor denotes floor rounding.

We look up the tables to retrieve corresponding embedding with the index and sum them up to obtain the position encoding of

5 Downsample and Upsample Layers

For the upsample layer, as shown in Fig. 6 (b), the decoder features xs′\mathbf{x}^{\prime}_{s} are firstly projected by a Pre-LN linear layer. We perform interpolation between current xyz coordinates ps\mathbf{p}_{s} and the previous ones ps−1\mathbf{p}_{s-1}. The encoder point features in the previous stage xs−1\mathbf{x}_{s-1} go through a Pre-LN linear layer. Finally, we sum them up to yield the next decoder features xs−1′\mathbf{x}^{\prime}_{s-1}.

Memory-efficient Implementation

In 2D Swin Transformer, it is easy to implement the window-based attention because the number of tokens is fixed in each window. Nevertheless, due to the irregular point arrangements in 3D, the number of the tokens in each window varies a lot. A simple solution is to pad the tokens in each window to the maximum token number kmaxk_{max} with dummy tokens, and then apply a masked self-attention. But this solution wastes much memory and computations.

Experiments

The main architecture is shown in Fig. 2. Both the xyz coordinates and rgb colors are used as inputs. We set the initial feature dimension and number of heads to 48 and 3 respectively, and they will double in each downsample layer. As for S3DIS, four stages are constructed with the block depths . In contrast, for ScanNetv2, we note that the point number is larger. So we add an extra downsample layer on top of the first-layer point embedding module. Then, the later four stages with block depths are added. So a total of five stages are constructed for ScanNetv2.

Implementation Detail.

For S3DIS, following previous work , we train for 76,50076,500 iterations with 4 RTX 2080Ti GPUs. The batch size is set to 88. Following common practice, the raw input points are firstly grid sampled with the grid size set to 0.040.04m. During training, the maximum input points number is set to 80,00080,000, and all extra ones are discarded if points number reaches this number. The window size is set to 0.160.16m initially, and it doubles after each downsample layer. The downsample scale for the stratified sampling strategy is set to 88. Unless otherwise specified, we use z-axis rotation, scale, jitter and drop color as data augmentation.

For ScanNetv2, we train for 600600 epochs with weight decay and batch size set to 0.10.1 and 88 respectively, and the grid size for grid sampling is set to 0.020.02m. At most 120,000120,000 points of a point cloud are fed into the network during training. The initial window size is set to 0.10.1m. And the downsample scale for the stratified sampling is set to 44. Except random jitter, the data augmentation is the same as that on S3DIS. The implementation details for ShapeNetPart and the datasets descriptions are given in the supplementary file.

2 Results

We make comparisons with recent state-of-the-art semantic segmentation methods. Tables 1 and 2 show the results on S3DIS and ScanNetv2 datasets. Our method achieves state-of-the-art performance on both challenging datasets. On S3DIS, ours outperforms others significantly, even higher than Point Transformer by 1.6%1.6\% mIoU. On ScanNetv2, the validation mIoU of our method surpasses others including voxel-based methods, with a gap of 2.1%2.1\% mIoU. On the test set, ours achieves slightly higher results than MinkowskiNet . The potential reason may be the points in ScanNetv2 are relatively sparse. So the loss of accurate position in voxelization is negligible for voxel-based methods. But on S3DIS where points are denser, our method outperforms MinkowskiNet with a huge gap, i.e., 6.6%6.6\% mIoU. Also, ours outperforms MinkowskiNet by 2.1%2.1\% mIoU on the validation set and is much more robust than MinkowskiNet when encountering various perturbations in testing, as shown in Table 9. Notably, it is the first time for the point-based methods to achieve higher performance compared with voxel-based methods on ScanNetv2.

Also, in Table 3, to show the generalization ability, we also make comparison on ShapeNetPart for the task of part segmentation. Our method outperforms previous ones and achieves new state of the art in terms of both category mIoU and instance mIoU. Although the instance mIoU of ours is comparable to Point Transformer, ours outperforms Point Transformer by a large margin in category mIoU.

3 Ablation Study

We conduct extensive ablation studies to verify the effectiveness of each component in our method, and show results in Table 4. To make our conclusions more convincing, we make evaluations on both S3DIS and ScanNetv2 datasets. From Exp.I to V, we add one component each time. Also, from Exp.VI to VIII, we make double verification by removing each component from the final model, i.e., Exp.V.

In Table 4, comparing Exp.IV and V, we notice that with the stratified strategy, the model improves with 1.9%1.9\% mIoU on S3DIS and 1.2%1.2\% mIoU on ScanNetv2. Combining the visualizations in Fig. 1, we note that the stratified strategy is able to enlarge the effective receptive field and boost the performance. Besides, we also show the effect when setting different downsample scales, i.e., 44, 88 and 1616, in the supplementary file.

First-layer Point Embedding.

We compare Exp.I with II, and find the model improves by a large margin with first-layer point embedding. Also, we compare Exp.VI and V, where the model gets 2.0%2.0\% mIoU gain on S3DIS and 4.0%4.0\% mIoU gain on ScanNetv2 with the equipment of first-layer point embedding. This minor modification in the architecture brings considerable benefit.

To further explore the role of local aggregation in first-layer point embedding, we compare different ways of local aggregation with linear projection in Table 5. Obviously, all listed local aggregation methods are better than linear projection for the first-layer point embedding.

Contextual Relative Position Encoding.

From Exp.III to IV, the performance increases by 2.9%2.9\% mIoU on S3DIS and 1.9%1.9\% mIoU on ScanNetv2 after using cRPE. Moreover, when also using the stratified Transformer, the model still improves with 4.0%4.0\% mIoU gain on S3DIS and 2.3%2.3\% gain on ScanNetv2 equipped with cRPE, through the comparison between Exp.VIII and V.

Further, we testify the contribution of applying cRPE on each of the query, key or value features. Table 6 shows that applying cRPE in either feature can make improvement. When applying cRPE on query, key and value simultaneously, the model achieves the best performance.

In addition, we compare our approach with the MLP-based method as mentioned in Sec. 5. As shown in Table 6, we find the MLP-based method (the first column) actually makes no difference with the model without any position encoding (the second column). Combining the visualization in Fig. 5, we conclude that the relative position information purely based on xyz coordinates is not helpful, since input point features to the network have already incorporated the xyz coordinates. In contrast, cRPE is based on both xyz coordinates and contextual features.

Shifted Window.

Shifted window is adopted to complement information interaction across windows. In Table 7, we compare the models w/ and w/o shifted window for both our vanilla version and Stratified Transformer on S3DIS. Evidently, shifted window is effective in our framework. Moreover, even without shifted window, Stratified Transformer still yields higher performance, i.e., 70.1%70.1\% mIoU, compared to the vanilla version. Also, shifting on both original and large windows is beneficial.

Data Augmentation.

Data augmentation plays an important role in training Transformer-based network. It is also the case in our framework as shown in Exp.V and VII as well as Exp.II and III. We also investigate the contribution of each augmentation in Table 8.

4 Robustness Study

To show the anti-interference ability of our model, we measure the robustness by applying a variety of perturbations in testing. Following , we make evaluations in aspects of permutation, rotation, shift, scale and jitter. As shown in Table 9, our method is extremely robust to various perturbations, while previous methods fluctuate drastically under these scenarios. It is notable that ours performs even better (+0.63%+0.63\% mIoU) with 90∘90^{\circ} z-axis rotation.

Although Point Transformer also employs the self-attention mechanism, it yields limited robustness. A potential reason may be Point Transformer uses special operator designs such as “vector self-attention” and “subtraction relation”, rather than standard multi-head self-attention.

5 Visual Comparison

In Fig. 8, we visually compare Point Transformer, the baseline model and ours. It clearly shows the superiority of our method. Due to the awareness of long-range contexts, our method is able to recognize the objects highlighted with yellow box, while others fail.

Conclusion

We propose Stratified Transformer and achieve state-of-the-art results. The stratified strategy significantly enlarges the effective receptive field. Also, first-layer point embedding and an effective contextual relative position encoding are put forward. Our work answers two questions. First, it is possible to build direct long-range dependencies at low computational costs and yield higher performance. Second, standard Transformer can be applied to 3D point cloud with strong generalization ability and powerful performance.

Acknowledgements

The work is supported in part by Hong Kong Research Grant Council - Early Career Scheme (Grant No. 27209621), HKU Startup Fund, HKU Seed Fund for Basic Research, and SmartMore donation fund.

References