Set Transformer: A Framework for Attention-based Permutation-Invariant Neural Networks

Juho Lee, Yoonho Lee, Jungtaek Kim, Adam R. Kosiorek, Seungjin Choi, Yee Whye Teh

Proofs

The decoder of a Set Transformer, given enough nodes, can express any element-wise function of the form (1n∑i=1nzip)1p\left(\frac{1}{n}\sum^{n}_{i=1}z_{i}^{p}\right)^{\frac{1}{p}}.

We first note that we can view the decoder as the composition of functions

We focus on HH in (2). Since feed-forward networks are universal function approximators at the limit of infinite nodes, let the feed-forward layers in front and back of the MAB encode the element-wise functions z→zpz\rightarrow z^{p} and z→z1pz\rightarrow z^{\frac{1}{p}}, respectively. We let h=dh=d, so the number of heads is the same as the dimensionality of the inputs, and each head is one-dimensional. Let the projection matrices in multi-head attention (WjQ,WjK,WjVW_{j}^{Q},W_{j}^{K},W_{j}^{V}) represent projections onto the jth dimension and the output matrix (WOW^{O}) the identity matrix. Since the mean operator is a special case of dot-product attention, by simple composition, we see that an MAB can express any dimension-wise function of the form

A PMA, given enough nodes, can express sum pooling (∑i=1nzi)\left(\sum^{n}_{i=1}z_{i}\right).

Set the seed ss to a zero vector and let ω(⋅)=1+f(⋅)\omega(\cdot)=1+f(\cdot), where ff is any activation function such that f(0)=0f(0)=0. The identiy, sigmoid, or relu functions are suitable choices for ff. The output of the multihead attention is then simply a sum of the values, which is ZZ in this case. ∎

We additionally have the following universality theorem for pooling architectures:

The Set Transformer is a universal function approximator in the space of permutation invariant functions.

While this proof required us to ignore the pairwise interaction terms inside the SABs and ISABs to prove that Set Transformers are universal function approximators, our experiments indicated that self-attention in the encoder was crucial for good performance.

Experiment Details

We show the detailed architectures used for the experiments in Table 1. We trained all networks using the Adam optimizer (Kingma2015) with a constant learning rate of 10−310^{-3} and a batch size of 128 for 20,000 batches, after which loss converged for all architectures.

2 Counting Unique Characters

The task generation procedure is as follows. We first sample a set size nn uniformly from the set of integers {6,…,10}\{6,\ldots,10\}. We then sample the number of characters cc uniformly from {1,…,n}\{1,\ldots,n\}. We sample cc characters from the training set of characters, and randomly sample instances of each character so that the total number of instances sums to nn and each set of characters has at least one instance in the resulting set.

The loss function we optimize, as previously mentioned, is the log likelihood log⁡p(x∣γ)=xlog⁡(γ)−γ−log⁡(x!)\log p(x|\gamma)=x\log(\gamma)-\gamma-\log(x!). We chose this loss function over mean squared error or mean absolute error because it seemed like the more logical choice when trying to make a real number match a target integer. Early experiments showed that directly optimizing for mean absolute error had roughly the same result as optimizing γ\gamma in this way and measuring ∣γ−x∣|\gamma-x|. We train using the Adam optimizer with a constant learning rate of 10−410^{-4} for 200,000 batches each with batch size 32.

3 Solving maximum likelihood problems for mixture of Gaussians

We generated the datasets according to the following generative process.

Table 4 summarizes the architectures used for the experiments. For all architectures, at each training step, we generate 10 random datasets according to the above generative process, and updated the parameters via Adam optimizer with initial learning rate 10−310^{-3}. We trained all the algorithms for 50k50k steps, and decayed the learning rate to 10−410^{-4} after 35k35k steps. Table 5 summarizes the detailed results with various number of inducing points in the ISAB. Figure LABEL:fig:synthetic_clustering shows the actual clustering results based on the predicted parameters.

3.2 2D Synthetic Mixtures of Gaussians Experiment on Large-scale Data

3.3 Details for CIFAR-100 amortized clutering experiment

We pretrained VGG net (Simonyan2014) with CIFAR-100, and obtained the test accuracy 68.54%. Then, we extracted feature vectors of 50k training images of CIFAR-100 from the 512-dimensional hidden layers of the VGG net (the layer just before the last layer). Given these feature vectors, the generative process of datasets is as follows.

Uniformly sample four classes among 100 classes.

Uniformly sample nn data points among four sampled classes.

Table 7 summarizes the architectures used for the experiments. For all architectures, at each training step, we generate 10 random datasets according to the above generative process, and updated the parameters via Adam optimizer with initial learning rate 10−410^{-4}. We trained all the algorithms for 50k50k steps, and decayed the learning rate to 10−510^{-5} after 35k35k steps. Table 8 summarizes the detailed results with various number of inducing points in the ISAB.

4 Set Anomaly Detection

Table 9 describes the architecture for meta set anomaly experiments. We trained all models via Adam optimizer with learning rate 10−410^{-4} and exponential decay of learning rate for 1,000 iterations. 1,000 datasets subsampled from CelebA dataset (see Figure LABEL:fig:celeba_figures) are used to train and test all the methods. We split 800 training datasets and 200 test datasets for the subsampled datasets.

5 Point Cloud Classification

We used the ModelNet40 dataset for our point cloud classification experiments. This dataset consists of a three-dimensional representation of 9,843 training and 2,468 test data which each belong to one of 40 object classes. As input to our architectures, we produce point clouds with n=100,1000,5000n=100,1000,5000 points each (each point is represented by (x,y,z)(x,y,z) coordinates). For generalization, we randomly rotate and scale each set during training.

Additional Experiments