PRNet: Self-Supervised Learning for Partial-to-Partial Registration
Yue Wang, Justin M. Solomon
Introduction
Registration is the problem of predicting a rigid motion aligning one point cloud to another. Algorithms for this task have steadily improved, using machinery from vision, graphics, and optimization. These methods, however, are usually orders of magnitude slower than “vanilla” Iterative Closest Point (ICP), and some have hyperparameters that must be tuned case-by-case. The trade-off between efficiency and effectiveness is steep, reducing generalizability and/or practicality.
Recently, PointNetLK and Deep Closest Point (DCP) show that learning-based registration can be faster and more robust than classical methods, even when trained on different datasets. These methods, however, cannot handle partial-to-partial registration, and their one-shot constructions preclude refinement of the predicted alignment.
We introduce the Partial Registration Network (PRNet), a sequential decision-making framework designed to solve a broad class of registration problems. Like ICP, our method is designed to be applied iteratively, enabling coarse-to-fine refinement of an initial registration estimate. A critical new component of our framework is a keypoint detection sub-module, which identifies points that match in the input point clouds based on co-contextual information. Partial-to-partial point cloud registration then boils down to detecting keypoints the two point clouds have in common, matching these keypoints to one another, and solving the Procrustes problem.
Since PRNet is designed to be applied iteratively, we use Gumbel–Softmax with a straight-through gradient estimator to sample keypoint correspondences. This new architecture and learning procedure modulates the sharpness of the matching; distant point clouds given to PRNet can be coarsely matched using a diffuse (fuzzy) matching, while the final refinement iterations prefer sharper maps. Rather than introducing another hyperparameter, PRNet uses a sub-network to predict the temperature of the Gumbel–Softmax correspondence, which can be cast as a simplified version of the actor-critic method. That is, PRNet learns to modulate the level of map sharpness each time it is applied.
We train and test PRNet on ModelNet40 and on real data. We visualize the keypoints and correspondences for shapes from the same or different categories. We transfer the learned representations to shape classification using a linear SVM, achieving comparable performance to state-of-the-art supervised methods on ModelNet40.
We summarize our key contributions as follows:
We present the Partial Registration Network (PRNet), which enables partial-to-partial point cloud registration using deep networks with state-of-the-art performance.
We use Gumbel–Softmax with straight-through gradient estimation to obtain a sharp and near-differentiable mapping function.
We design an actor-critic closest point module to modulate the sharpness of the correspondence using an action network and a value network. This module predicts more accurate rigid transformations than differentiable soft correspondence methods with fixed parameters.
We show registration is a useful proxy task to learn representations for 3D shapes. Our representations can be transferred to other tasks, including keypoint detection, correspondence prediction, and shape classification.
We release our code to facilitate reproducibility and future research. https://github.com/WangYueFt/prnet
Related Work
Rigid Registration. ICP and variants have been widely used for registration. Recently, probabilistic models have been proposed to handle uncertainty and partiality. Another trend is to improve the optimization: applies Levenberg–-Marquardt to the ICP objective, while global methods seek a solution using branch-and-bound , Riemannian optimization , convex relaxation , mixed-integer programming , and semidefinite programming .
Learning on Point Clouds and 3D Shapes. Deep Sets and PointNet pioneered deep learning on point sets, a challenge problem in learning and vision. These methods take coordinates as input, embed them to high-dimensional space using shared multilayer perceptrons (MLPs), and use a symmetric function (e.g., or ) to aggregate features. Follow-up works incorporate local information, including PointNet++ , DGCNN , PointCNN , and PCNN . Another branch of 3D learning designs convolution-like operations for shapes or applies graph convolutional networks (GCNs) to triangle meshes , exemplifying architectures on non-Euclidean data termed geometric deep learning . Other works, including SPLATNet , SplineCNN , KPConv , and GWCNN , transform 3D shapes to regular grids for feature learning.
Keypoints and Correspondence. Correspondence and registration are dual tasks. Correspondence is the approach while registration is the output, or vice versa. Countless efforts tackle the correspondence problem, either at the point-to-point or part-to-part level. Due to the complexity of point-to-point correspondence matrices and possible permutations, most methods (e.g., ) compute a sparse set of correspondences and extend them to dense maps, often with bijectivity as an assumption or regularizer. Other efforts use more exotic representations of correspondences. For example, functional maps generalize to mappings between functions on shapes rather than points on shapes, expressing a map as a linear operator in the Laplace–Beltrami eigenbasis. Mathematical methods like functional maps can be made ‘deep’ using priors learned from data: Deep functional maps learn descriptors rather than designing them by hand.
For partial-to-partial registration, we cannot compute bijective correspondences, invalidating many past representations. Instead, keypoint detection is more secure. To extract a sparser representation, KeyPointNet uses registration and multiview consistency as supervision to learn a keypoint detector on 2D images; our method performs keypoint detection on point clouds. In contrast to our model, which learns correspondences from registration, uses correspondence prediction as the training objective to learn how to segment parts. In particular, it utilizes PointNet++ to product point-wise features, generates matching using a correspondence proposal module, and finally trains the pipeline with ground-truth correspondences.
Self-supervised Learning. Humans learn knowledge not only from teachers but also by predicting and reasoning about unlabeled information. Inspired by this observation, self-supervised learning usually involves predicting part of an input from another part , solving one task using features learned from another task and/or enforcing consistency from different views/modalities . Self-supervised pretraining is an effective way to transfer knowledge learned from massive unlabeled data to tasks where labeled data is limited. For example, BERT surpasses state-of-the-art in natural language processing by learning from contextual information. ImageNet Pretrain commonly provides initialization for vision tasks. Video-audio joint analysis utilizes modality consistency to learn representations. Our method is also self-supervised, in the sense that no labeled data is needed.
Actor–Critic Methods. Many recent works can be counted as actor–critic methods, including deep reinforcement learning , generative modeling , and sequence generation . These methods generally involve two functions: taking actions and estimating values. The predicted values can be used to improve the actions while the values are collected when the models interact with environment. PRNet uses a sub-module (value head) to predict the level of granularity at which we should map two shapes. The value adjusts the temperature of Gumbel–Softmax in the action head.
Method
We establish preliminaries about the rigid alignment problem and related algorithms in §3.1; then, we present PRNet in §3.2. For ease of comparison to previous work, we use the same notation as .
where and are obtained using the singular value decomposition (SVD) , with . In this expression, centroids of and are defined as and respectively.
We can understand ICP and the more recent learning-based DCP method as providing different choices of :
Iterative Closest Point. ICP chooses to minimize (1) with fixed, yielding:
ICP approaches a fixed point by alternating between (2) and (3); each step decreases the objective (1). Since (1) is non-convex, however, there is no guarantee that ICP reaches a global optimum.
Deep Closest Point. DCP uses deep networks to learn . In this method, and are embedded using learned functions and defined by a Siamese DGCNN ; these lifted point clouds are optionally contextualized by a Transformer module , yielding embeddings and . The mapping is then
This formula is applied in one shot followed by (2) to obtain the rigid alignment. The loss used to train this pipeline is mean-squared error (MSE) between ground-truth rigid motion from synthetically-rotated point clouds and prediction; the network is trained end-to-end.
2 Partial Registration Network
DCP is a one-shot algorithm, in that a single pass through the network determines the output for each prediction task. Analogously to ICP, PRNet is designed to be iterative; multiple passes of a point cloud through PRNet refine the alignment. The steps of PRNet, illustrated in Figure 1, are as follows:
take as input point clouds and ;
detect keypoints of and ;
predict a mapping from keypoints of to keypoints of ;
predict a rigid transformation aligning to based on the keypoints and map;
transform using the obtained transformation;
return to 1 using the pair as input.
When predicting a mapping from keypoints in to keypoints in , PRNet uses Gumbel–Softmax to sample a matching matrix, which is sharper than (4) and approximately differentiable. It has a value network to predict a temperature for Gumbel–Softmax, so that the whole framework can be seen as an actor-critic method. We present details of and justifications behind the design below.
Notation. Denote by the rigid motion of to align to after applications of PRNet; and are initial input shapes. We will use to denote the -th rigid motion predicted by PRNet for the input pair .
Since our training pairs are synthetically generated, before applying PRNet we know the ground-truth aligning to . From these values, during training we can compute “local” ground-truth on-the-fly, which maps the current to the best alignment:
We use to denote the mapping function in -th step.
Synthesizing the notation above, is given by
In this equation, and are computed using (2) from , , and .
Keypoint Detection. For partial-to-partial registration, usually and only subsets of and match to one another. To detect these mutually-shared patches, we design a simple yet efficient keypoint detection module based on the observation that the norms of features tend to indicate whether a point is important.
Using and to denote the keypoints for and , we take
By aligning only the keypoints, we remove irrelevant points from the two input clouds that are not shared in the partial correspondence. In particular, we can now solve the Procrustes problem that matches keypoints of and . We show in §4.3 that although we do not provide explicit supervision, PRNet still learns how to detect keypoints reasonably.
Gumbel–Softmax Sampler. One key observation in ICP and DCP is that (3) usually is not differentiable with respect to the map but by definition yields a sharp correspondence between the points in and the points in . In contrast, the smooth function (4) in DCP is differentiable, but in exchange for this differentiability the mapping is blurred. We desire the best of both worlds: A potentially sharp mapping function that admits backpropagation.
To that end, we use Gumbel–Softmax to sample a matching matrix. Using a straight-through gradient estimator, this module is approximately differentiable. In particular, the Gumbel–Softmax mapping function is given by
Actor-Critic Closest Point (ACP). The mapping functions (4) and (10) have fixed “temperatures,” that is, there is no control over the sharpness of the mapping matrix . In PRNet, we wish to adapt the sharpness of the map based on the alignment of the two shapes. In particular, for low values of (the initial iterations of alignment) we may satisfied with high-entropy approximate matchings that obtain a coarse alignment; later during iterative evaluations, we can sharpen the map to align individual pairs of points.
To make this intuition compatible with PRNet’s learning-based architecture, we add a parameter to (10) to yield a generalized Gumbel–Softmax matching matrix:
When is large, the map matrix is smoothed out; as the map approaches a binary matrix.
Loss Function. The final loss is the summation of several terms , indexed by the number of passes through PRNet for the input pair. consists of three terms: a rigid motion loss , a cycle consistency loss , and a global feature alignment loss . We also introduce a discount factor to promote alignment within the first few passes through PRNet; during training we pass each input pair through PRNet times.
Equation (5) gives the “localized” ground truth values for . Denoting the rigid motion from to in step as the cycle consistency loss is
Our last loss term is a global feature alignment loss, which enforces alignment of global features and . Mathematically, the global feature alignment loss is
This global feature alignment loss also provides signal for determining . When two shapes are close in global feature space, should be small, yielding a sharp matching matrix; when two shapes are far from each other, increases and the map is blurry.
Experiments
Our experiments are divided into four parts. First, we show performance of PRNet on a partial-to-partial registration task on synthetic data in §4.1. Then, we show PRNet can generalize to real data in §4.2. Third, we visualize the keypoints and correspondences predicted by PRNet in §4.3. Finally, we show a linear SVM trained on representations learned by PRNet can achieve comparable results to supervised learning methods in §4.4.
We evaluate partial-to-partial registration on ModelNet40 . There are 12,311 CAD models spanning 40 object categories, split to 9,843 for training and 2,468 for testing. Point clouds are sampled from the CAD models by farthest-point sampling on the surface. During training, a point cloud with 1024 points is sampled. Along each axis, we randomly draw a rigid transformation; the rotation along each axis is sampled in and translation is in . We apply the rigid transformation to , leading to . We simulate partial scans of and by randomly placing a point in space and computing its 768 nearest neighbors in and respectively.
We measure mean squared error (MSE), root mean squared error (RMSE), mean absolute error (MAE), and coefficient of determination (R2). Angular measurements are in units of degrees. MSE, RMSE and MAE should be zero while R2 should be one if the rigid alignment is perfect. We compare our model to ICP, Go-ICP , Fast Global Registration (FGR) , and DCP .
Figure 1 shows the architecture of ACP. We use DGCNN with 5 dynamic EdgeConv layers and a Transformer to learn co-contextual representations of and . The number of filters in each layer of DGCNN are . In the Transformer, only one encoder and one decoder with 4-head attention are used. The embedding dimension is 1024. We train the network for 100 epochs using Adam . The initial learning rate is 0.001 and is divided by 10 at epochs 30, 60, and 80.
Partial-to-Partial Registration on Unseen Objects. We first evaluate on the ModelNet40 train/test split. We train on 9,843 training objects and test on 2,468 testing objects. Table 1 shows performance. Our method outperforms its counterparts in all metrics.
Partial-to-Partial Registration on Unseen Categories. We follow the same testing protocol as to compare the generalizability of different models. ModelNet40 is split evenly by category into training and testing sets. PRNet and DCP are trained on the first 20 categories, and then all methods are tested on the held-out categories. Table 2 shows PRNet behaves more strongly than others. To further test generalizability, we train it on ShapeNetCore dataset and test on ModelNet40 held-out categories. ShapeNetCore has 57,448 objects, and we do the same preprocessing as on ModelNet40. The last row in Table 2, denoted as PRNet (Ours*), surprisingly shows PRNet performs much better than when trained on ModelNet40. This supports the intuition that data-driven approaches work better with more data.
Partial-to-Partial Registration on Unseen Objects with Gaussian Noise. We further test robustness to noise. The same preprocessing is done as in the first experiment, except that noise independently sampled from and clipped to is added to each point. As in Table 3, learning-based methods, including DCP and PRNet, are more robust. In particular, PRNet exhibits stronger performance and is even comparable to the noise-free version in Table 1.
2 Partial-to-Partial on Real Data
We test our model on the Stanford Bunny dataset . Since the dataset only has 10 real scans, we fine tune the model used in Table 1 for 10 epochs with learning rate 0.0001. For each scan, we generate 100 training examples by randomly transforming the scan in the same way as we do in §4.1. This training procedure can be viewed as inference time fine-tuning, in contrast to optimization-based methods that perform one-time inference for each test case. Figure 2 shows the results. We further test our model on more scans from Stanford 3D Scanning Repository using a similar methodology; Figure 3 shows the registration results.
3 Keypoints and Correspondences
We visualize keypoints on several objects in Figure 4 and correspondences in Figure 6. The model detects keypoints and correspondences on partially observable objects. We overlay the keypoints on top of the fully observable objects. Also, as shown in Figure 5, the keypoints are consistent across different views.
4 Transfer to Classification
Conclusion
PRNet tackles a general partial-to-partial registration problem, leveraging self-supervised learning to learn geometric priors directly from data. The success of PRNet verifies the sensibility of applying learning to partial matching as well as the specific choice of Gumbel–Softmax, which we hope can inspire additional work linking discrete optimization to deep learning. PRNet is also a reinforcement learning-like framework; this connection between registration and reinforcement learning may provide inspiration for additional interdisciplinary research related to rigid/non-rigid registration.
Our experiments suggest several avenues for future work. For example, as shown in Figure 6, the matchings computed by PRNet are not bijective, evident e.g. in the point clouds of cars and chairs. One possible extension of our work to address this issue is to use Gumbel–Sinkhorn to encourage bijectivity. Improving the efficiency of PRNet when applied to real scans also will be extremely valuable. As described in §4.2, PRNet currently requires inference-time fine-tuning on real scans to learn useful data-dependent representations; this makes PRNet slow during inference. Seeking universal representations that generalize over broader sets of registration tasks will improve the speed and generalizability of learning-based registration. Another possibility for future work is to improve the scalability of PRNet to deal with large-scale real scans captured by LiDAR.
Finally, we hope to find more applications of PRNet beyond the use cases we have shown in the paper. A key direction bridging PRNet to applications will involve incorporating our method into SLAM or structure-from-motion can demonstrate its value for robotics applications and robustness to realistic species of noise. Additionally, we can test the effectiveness of PRNet for registration problems in medical imaging and/or high-energy particle physics.
Acknowledgements
The authors acknowledge the generous support of Army Research Office grant W911NF1710068, Air Force Office of Scientific Research award FA9550-19-1-031, of National Science Foundation grant IIS-1838071, from an Amazon Research Award, from the MIT-IBM Watson AI Laboratory, from the Toyota-CSAIL Joint Research Center, from a gift from Adobe Systems, and from the Skoltech-MIT Next Generation Program. Any opinions, findings, and conclusions or recommendations expressed in this material are those of the authors and do not necessarily reflect the views of these organizations. The authors also thank members of MIT Geometric Data Processing group for helpful discussion and feedback on the paper.
References
Supplementary
We provide more details of PRNet in this section.
A shared DGCNN is to use extract embeddings for and separately. The number of filters per layer are . We use BatchNorm and LeakyReLU after each MLP in the EdgeConv layer. The local aggregation function of -nn graph is , and there is no global aggregation function used in DGCNN. and denote the representations learned by DGCNN.
After DGCNN, and are fed into the Transformer. The Transformer is an asymmetric function that learns co-contextual representations and . Transformer has only one encoder and one decoder. 4-head self-attention is used in encoder and decoder. LayerNorm, instead of BatchNorm, is used in the Transformer. Unlike the original implementation of Transformer, we do not use Dropout. For detailed presentation of Transformer, we refer readers to the tutorial.http://nlp.seas.harvard.edu/2018/04/03/attention.html
There are two heads on top of the representations and : a action head consisting of Gumbel-Softmax and SVD; a value head to predict a for Gumbel-Softmax in the action head. The value head is parameterized by a 4-layer MLPs. The number of filters are . BatchNorm and ReLU are used after each linear layer in the MLPs.
Training Protocol.
We train the model for 100 epochs. At epochs 30, 60, and 80, we divide the learning rate by 10; it is initially 0.001. Each training pair and is passed through PRNet three times iteratively (the rigid alignment of is updated three times). The final rigid transformation is the combination of these three local rigid transformations. for cycle consistency loss and for feature alignment loss are both 0.1. The weight decay used is . The number of keypoints is 512 on training. For visualization purposes, however, we show 64 keypoints in Figure 4, Figure 5, and Figure 6.
Our model is trained on a Google Cloud GPU instance with 4 Tesla V100 GPUs and takes 10 hours to complete.
As for DCP-v2, we take the implementation from the authors’ released code https://github.com/WangYueFt/dcp and train it as they suggest.
Choices of λ𝜆\lambda.
We compare to alternative choices of ways to determine : (1) fixing manually; (2) annealing to near 0 as the training going; (3) including as a variable during training. We train the PRNet in the same way for each option, except the choice of is different. Table 5 verifies our choice of strategies for computing .
To understand the effectiveness of each part, we conduct additional experiments in Table 6; to save space, we only show MAE and R2. (a) First, we consider alternatives to keypoint selection: in the first alternative, the two sets of keypoints are chosen independently and randomly on the two surfaces ( and ); in the second alternative, we use centrality to choose keypoints, keeping the points whose average distance (in feature space) to the rest in the point cloud is minimal. Empirically, the norm used in our pipeline to select keypoints outperforms others. (b) Second, we compare our method to others on full point clouds. In this experiment, 768 points are sampled from each point cloud to cover the full shape using farthest-point sampling. In the full point cloud setting, PRNet still outperforms others. (c) Third, we verify our choice of discount factor ; small large discount factors encourage alignment within the first few passes through PRNet while large discount factors promote longer-term return. (d) Fourth, we test the choice of number of keypoints: the model achieves surprisingly good performance even with 64 keypoints, but performance drops significantly when . (e) Fifth, we test its robustness to missing data. The missing data ratio in original partial-to-partial experiment is 25%; we further test with 50% and 75%. This test shows that with 75% points missing, the method still achieves reasonable performance, even compared to other methods tested with only 25% points missing. (f) Finally, we test the model robustness to noise level. Noise is sampled from . The model is trained with and tested with . Even with , the model still performs reasonably well.
Efficiency.
We benchmark the inference time of different methods on a desktop computer with an Intel 16-core CPU, an Nvidia GTX 1080 Ti GPU, and 128G memory. Table 7 shows learning based methods (on GPUs) are faster than non-learning based counterparts (on CPUs). PRNet is on a par with PointNetLK while being slower than DCP.
More figures of keypoints and correspondences.
In Figure 7 and Figure 8, we show more visualizations of keypoints and correspondences for different pairs of objects.