Conditional Set Generation with Transformers
Adam R Kosiorek, Hyunjik Kim, Danilo J Rezende
Introduction
It is natural to reason about a group of objects as a set. Therefore many machine learning tasks involving predicting objects or their properties can be cast as a set prediction problem. These predictions are usually conditioned on some input feature that can take the form of a vector, a matrix or a set. Some examples include predicting future states for a group of molecules in a simulation Noé et al. (2020), object detection from images Carion et al. (2020) and generating correlated samples for sequential Monte Carlo in object tracking Zhu et al. (2020); Neiswanger et al. (2014). Elements of a set are unordered, which brings about two challenges that set prediction faces. First, the model must be permutation-equivariant; that is, the generation of a particular permutation of the set elements must be equally probable to any other permutation. Second, training a generative model for sets typically involves comparing a predicted set against a ground-truth set. Since the result of this comparison should not depend on the permutation of the elements of either set, the loss function used for training must be permutation-invariant. While it is possible to create a set model that violates either or both of these requirements, such a model has to learn to meet them, which is likely to result in lower performance.
Permutation equivariance imposes a constraint on the structure of the model Bloem-Reddy & Teh (2018; 2019), and therefore sets are often treated as an ordered collection of items, which allows using standard machine learning models. For example assuming that a set has a fixed size, we can treat it as a tensor and turn set-prediction into multivariate regression Achlioptas et al. (2018). If the ordering is fixed but the size is not, we can treat set prediction as a sequence prediction problem Vinyals et al. (2016). Both approaches require using permutation-invariant loss functions to allow the model to learn a deterministic ordering policy Eslami et al. (2016). However imposing such an ordering can lead to a pathology that is commonly referred to as the responsibility problem (Zhang et al., 2019; 2020); there exist points in the output set space where a small change in set space (as measured by a set loss) requires a large change in the generative model’s output. This can lead to sub-optimal performance, as shown in Zhang et al. (2020). Some approaches choose to learn the ordering of set elements Rezatofighi et al. (2018), but this also suffers the same problem as well as adding further complexity to the set prediction problem.
Recently, Zhang et al. (2019) introduced the Deep Set Prediction Network (dspn)—a model that generates sets in a permutation-equivariant manner using permutation-invariant loss functions. Dspn relies on the observation that the gradient of a permutation-invariant function is equivariant with respect to the permutation of its inputs, also noticed by Papamakarios et al. (2019). Dspn uses this insight to generate a set by gradient-descent on a learned loss function with respect to an initially-guessed set. Dspn has several limitations, however. The functional form of the update step is limited, as the gradient information is only used to tanslate the set elements. This, in turn, means that the method can be computationally costly: not only is the backward pass expensive, but many such passes might be needed to arrive at an accurate prediction.
In this paper we develop the Transformer Set Prediction Network (tspn), where we replace the gradient-based updates of dspn with a Transformer Vaswani et al. (2017), that is also permutation-equivariant and learns to jointly update the elements of the initial set. We make the following contributions:
We show that the set-cardinality learning method of Zhang et al. (2019) is prone to falling into local minima.
We thus introduce an alternative, principled method for learning set cardinality.
We learn a distribution over the elements of the initial set (as opposed to dspn that learns a fixed initial set). This allows one to directly generate sets with cardinality determined by the model (by sampling the correct number of points), and to dynamically change the size of the generated sets as test time.
We demonstrate that tspn outperforms dspn on conditional point-cloud generation and object-detection tasks.
We show that our model is not only more expressive than the dspn, but can also generalize at test-time to sets of vastly different cardinality than the sets encountered during training. We evaluate our model on auto-encoding set-mnist LeCun et al. (2010) and on object detection on clevr Johnson et al. (2017). We now proceed to describe dspn and our method tspn in detail, followed by experiments.
Permutation-Equivariant Set Generation
In dspn, the points are initialised randomly as model parameters and learned, hence the model assumes fixed set cardinality. For handling variable set sizes, each element of the predicted set is augmented with a presence variable , which are transformed by along with the and then thresholded to give the final prediction. The ground-truth sets are padded with all-zero vectors to the same size. Note, however, that this mechanism does not allow to extrapolate to set sizes beyond the maximum size encountered in training.
Dspn employs a permutation-invariant set encoder (e.g. deep-sets Zaheer et al. (2017), relation networks Santoro et al. (2017)) to produce a set embedding, and updates the initial prediction using the gradients of an embedding loss, arriving at a permutation-equivariant final prediction.
2 Permutation-Invariant Loss
The model is trained using a set loss, i.e. a permutation-invariant loss. Common choices are the Chamfer loss or the Hungarian loss:
where is a permutation in the space of all possible permutations and can be any distance or loss function defined on pairs of set points. Note that the computational complexity of the Chamfer loss is for sets of size , whereas the Hungarian loss is —it uses the Hungarian algorithm to compute the optimal permutation, whose complexity is Bayati et al. (2008). Hence the Chamfer loss is suitable for larger sets, and the Hungarian loss for smaller sets. For , Zhang et al. (2019) use the Huber loss defined as .
Recall that in the implementation of dspn, the ground truth set is padded with zero vectors so that all sets have the same size. Padding a set to a fixed size with constant elements turns it into a multiset . A Multiset is a set that contains repeated elements, that can be represented by an augmented set where each unique element is paired with its multiplicity. If we use a multiset in its default form (i. e. with repeated elements) as the ground-truth in the Chamfer loss, then it is enough for the model to predict a set containing exactly one element equal to the repeated element of the set in order to account for all its repetitions in the first term of Equation 4. The remaining superfluous elements predicted by the model can match any other element in without increasing the second term of Equation 4. This implies that padding a ground-truth set of size with constant elements creates predictions that are all optimal and hence indistinguishable under the Chamfer loss. These predictions have a set size that varies from to , and hence the model is likely to fail to learn the correct set cardinality—an effect clearly visible in our experiments, c. f. Section 4.1 and Table 1.
Size-Conditioned Set Generation with Transformers
While training tspn, we use the ground-truth set-cardinality to instantiate the initial set, and we separately train the mlp by minimizing categorical cross-entropy with the ground-truth set sizes. The cardinality-mlp is used only at test-time. Note that, in contrast to dspn, tspn does not require any additional regularization terms applied to its representations.
We describe our work in relation to dspn Zhang et al. (2019), but we note that there are concurrent works that share the ideas presented here. Both detr Carion et al. (2020) and Slot Attention Locatello et al. (2020) use a variant of the transformer Vaswani et al. (2017) for predicting a set of object properties. Detr uses an object-specific query initialization and is, therefore, not equivariant to permutations, similarly to Zhang et al. (2019). Slot Attention is perhaps the most similar to our work—transforms a randomly sampled point-cloud, same as tspn, but it uses attention normalized along the query axis instead the key axis.
Experiments
We apply tspn and our implementations of dspn and a size-conditioned dspn (c-dspn) to two tasks: point-cloud prediction on set-mnist and object detection on clevr. Point-cloud prediction is an autoencoding task, where we use a set encoder to produce vector embeddings of point clouds of varying sizes, and explore the use of the above set prediction networks for reconstructing the point clouds conditioned on these embeddings. Object detection, instead, requires predicting sets of bounding boxes conditioned on image features. While generating point-clouds requires predicting large numbers of points, which are often assumed to be conditionally independent of each other, detecting objects typically requires generating much smaller sets. Due to possible overlaps between neighbouring bounding boxes, however, it is essential to take relations between different objects into account. We implemented all models in jax Bradbury et al. (2018) and haiku Hennigan et al. (2020) and run experiments on a single Tesla V100 GPU. Models are optimized with adam Kingma & Ba (2015) with default parameters. We used batch size and trained all models for epochs.
We convert mnist into point clouds by thresholding pixel values with the mean value in the training set; then normalize coordinates of remaining points to be in .
We use the same hyperparameters for dspn as reported in Zhang et al. (2019). c-dspn uses the same settings, except instead of presence variables it uses an mlp with one hidden layer of size to predict the set size, as used in tspn. Tspn uses a three-layer Transformer, with parameters shared between layersThis gives 383k parameters compared to 190k for dspn, which shares parameters between its encoder and gradient-based decoder. Not sharing parameters between layers does not improve results, but significantly increases parameter count. Sharing parameters and increasing layer width to match the number of parameters does not increase performance, either.. Each layer has hidden units and attention heads. The initial set is two-dimensional (same as the output), but it is linearly-projected into -dimensional space. The outputs of the transformer layers are kept at dimensions, too, and we project them back into dimensions after the last layer. Tspn uses the same input encoder as (c-)dspn, that is a two-layer deepset with fspool pooling layer Zhang et al. (2020). All models are trained using the Chamfer loss with learning rate .
Table 1 shows quantitative results. Conditioning on the set size does not improve the Chamfer loss for c-dspn but it does significantly improve the accuracy of set-size prediction. Further, replacing the decoder with the Transformer (Tspn) leads to a significant loss reduction with respect to dspn. Figure 3 shows inputs set and model reconstructions. Note that in our experiments, dspn almost always predicts the same number of points (), which is close to the maximum number of points we used (here, ). Interestingly, tspn performs very well while extrapolating to much bigger sets than the ones encountered in training: Figure 6 in Appendix A shows reconstructions, where we manually change the desired set size up to points. This is in contrast to c-dspn, whose performance decreases significantly when we require it to generate a set whose size differs only slightly from the input, c. f. Figure 6 in Appendix A. We conjecture that this is caused by how fspool handles sets of different sizes, which leads to incompatibility between embeddings of sets of different cardinality. This, in turn, causes the distances in the latent space to be ill-defined, and breaks the internal gradient-based optimization.
2 Object Detection on Clevr
Clevr images consist of up to 10 rendered objects on plain backgrounds. Following Zhang et al. (2019), we use the clevr dataset to test the efficacy of our models for object detection in a simple setting, which might pave the way for more advanced object-based inference in future works. We use the same hyperparameters for c-dspn as Zhang et al. (2019). Tspn uses four layers with four attention heads and neurons each, without parameter sharing between layers. We apply layer normalization Xiong et al. (2020) before attention as in Ba et al. (2016). The last transformer layer is followed by an mlp with a single hidden layer of All models use a resnet34 He et al. (2016) as the input encoder, and are trained with the Chamfer loss; dspn-based models are trained for epochs with learning rate ; tspn uses learning rate and is trained for epochs. Longer training for dspn-based models lead to overfititng and decreased validation performance. Note that dspn in Zhang et al. (2019) use the Hungarian loss (Equation 5), which leads to better results than using Chamfer. We report the average precision scores at different thresholds and the set size root-mean-square error in Table 2. Qualitative results are available in Figure 4. We see that tspn outperforms c-dspn and produces bounding boxes that are better aligned with ground-truth, although these results are worse than the ones reported in Zhang et al. (2019)—we expect improvements when using the Hungarian loss as well.
Conclusions
We introduced the Transformer Set Prediction Network (tspn)—a transformer-based model for conditional set prediction. Tspn infers the cardinality of the set, randomly samples an initial set of the desired size and applies a transformer to generate the final prediction. Set prediction in tspn is permutation-equivariant, and the model can be applied to any set-prediction tasks. Interesting directions include scaling the model to large-scale point-clouds and object detection (e. g. similar to Carion et al. (2020)), as well as turning this model into a generative model in either the vae or gan framework.
Acknowledgements
We would like to thank George Papamakarios, Karl Stelzner, Thomas Kipf, Teophane Weber and Yee Whye Teh for helpful discussions.