E(n) Equivariant Graph Neural Networks
Victor Garcia Satorras, Emiel Hoogeboom, Max Welling
Introduction
Although deep learning has largely replaced hand-crafted features, many advances are critically dependent on inductive biases in deep neural networks. An effective method to restrict neural networks to relevant functions is to exploit the symmetry of problems by enforcing equivariance with respect to transformations from a certain symmetry group. Notable examples are translation equivariance in Convolutional Neural Networks and permutation equivariance in Graph Neural Networks (Bruna et al., 2013; Defferrard et al., 2016; Kipf & Welling, 2016a).
Many problems exhibit 3D translation and rotation symmetries. Some examples are point clouds (Uy et al., 2019), 3D molecular structures (Ramakrishnan et al., 2014) or N-body particle simulations (Kipf et al., 2018). The group corresponding to these symmetries is named the Euclidean group: SE(3) or when reflections are included E(3). It is often desired that predictions on these tasks are either equivariant or invariant with respect to E(3) transformations.
Recently, various forms and methods to achieve E(3) or SE(3) equivariance have been proposed (Thomas et al., 2018; Fuchs et al., 2020; Finzi et al., 2020; Köhler et al., 2020). Many of these works achieve innovations in studying types of higher-order representations for intermediate network layers. However, the transformations for these higher-order representations require coefficients or approximations that can be expensive to compute. Additionally, in practice for many types of data the inputs and outputs are restricted to scalar values (for instance temperature or energy, referred to as type- in literature) and 3d vectors (for instance velocity or momentum, referred to as type- in literature).
We evaluate our method in modelling dynamical systems, representation learning in graph autoencoders and predicting molecular properties in the QM9 dataset. Our method reports the best or very competitive performance in all three experiments.
Background
In this section we introduce the relevant materials on equivariance and graph neural networks which will later complement the definition of our method.
Let be a set of transformations on for the abstract group . We say a function is equivariant to if there exists an equivalent transformation on its output space such that:
Permutation equivariance. Permuting the input results in the same permutation of the output where is a permutation on the row indexes.
2 Graph Neural Networks
Graph Neural Networks are permutation equivariant networks that operate on graph structured data (Bruna et al., 2013; Defferrard et al., 2016; Kipf & Welling, 2016a). Given a graph with nodes and edges we define a graph convolutional layer following notation from (Gilmer et al., 2017) as:
Equivariant Graph Neural Networks
Our Equivariant Graph Convolutional Layer (EGCL) takes as input the set of node embeddings , coordinate embeddings and edge information and outputs a transformation on and . Concisely: . The equations that define this layer are the following:
Notice the main differences between the above proposed method and the original Graph Neural Network from equation 2 are found in equations 3 and 4. In equation 3 we now input the relative squared distance between two coordinates into the edge operation . The embeddings , , and the edge attributes are also provided as input to the edge operation as in the GNN case. In our case the edge attributes will incorporate the edge values , but they can also include additional edge information.
Finally, equations 5 and 6 follow the same updates than standard GNNs. Equation 5 is the aggregation step, in this work we choose to aggregate messages from all other nodes , but we could limit the message exchange to a given neighborhood if desired in both equations 5 and 4. Equation 6 performs the node operation which takes as input the aggregated messages , the node emedding and outputs the updated node embedding .
2 Extending EGNNs for vector type representations
In this section we propose a slight modification to the presented method such that we explicitly keep track of the particle’s momentum. In some scenarios this can be useful not only to obtain an estimate of the particle’s velocity at every layer but also to provide an initial velocity value in those cases where it is not 0. We can include momentum to our proposed method by just replacing Equation 4 of our model with the following equation:
3 Inferring the edges
Given a point cloud or a set of nodes, we may not always be provided with an adjacency matrix. In those cases we can assume a fully connected graph where all nodes exchange messages with each other as done in Equation 5. This fully connected approach may not scale well to large point clouds where we may want to locally limit the exchange of messages to a neighborhood to avoid an overflow of information.
Similarly to (Serviansky et al., 2020; Kipf et al., 2018), we present a simple solution to infer the relations/edges of the graph in our model, even when they are not explicitly provided. Given a set of neighbors for each node , we can re-write the aggregation operation from our model (eq. 5) in the following way:
Related Work
Experiments
In a dynamical system a function defines the time dependence of a point or set of points in a geometrical space. Modelling these complex dynamics is crucial in a variety of applications such as control systems (Chua et al., 2018), model based dynamics in reinforcement learning (Nagabandi et al., 2018), and physical systems simulations (Grzeszczuk et al., 1998; Watters et al., 2017). In this experiment we forecast the positions for a set of particles which are modelled by simple interaction rules, yet can exhibit complex dynamics.
Similarly to (Fuchs et al., 2020), we extended the Charged Particles N-body experiment from (Kipf et al., 2018) to a 3 dimensional space. The system consists of 5 particles that carry a positive or negative charge and have a position and a velocity associated in 3-dimensional space. The system is controlled by physic rules: particles are attracted or repelled depending on their charges. This is an equivariant task since rotations and translations on the input set of particles result in the same transformations throughout the entire trajectory.
Implementation details: In this experiment we used the extension of our model that includes velocity from section 3.2. We input the position as the first layer coordinates of our model and the velocity as the initial velocity in Equation 7, the norms are also provided as features to through a linear mapping. The charges are input as edge attributes . The model outputs the last layer coordinates as the estimated positions. We compare our method to its non equivariant Graph Neural Network (GNN) cousin, and the equivariant methods: Radial Field (Köhler et al., 2019), Tensor Field Networks and the SE(3) Transformer. All algorithms are composed of 4 layers and have been trained under the same conditions, batch size 100, 10.000 epochs, Adam optimizer, the learning rate was tuned independently for each model. We used 64 features for the hidden layers in the Radial Field, the GNN and our EGNN. As non-linearity we used the Swish activation function (Ramachandran et al., 2017). For TFN and the SE(3) Transformer we swept over different number of vector types and features and chose those that provided the best performance. Further implementation details are provided in Appendix C.1. A Linear model that simply considers the motion equation is also included as a baseline. We also provide the average forward pass time in seconds for each of the models for a batch of 100 samples in a GTX 1080 Ti GPU.
Results As shown in Table 2 our model significantly outperforms the other equivariant and non-equivariant alternatives while still being efficient in terms of running time. It reduces the error with respect to the second best performing method by a . In addition it doesn’t require the computation of spherical harmonics which makes it more time efficient than Tensor Field Networks and the SE(3) Transformer.
2 Graph Autoencoder
A Graph Autoencoder can learn unsupervised representations of graphs in a continuous latent space (Kipf & Welling, 2016b; Simonovsky & Komodakis, 2018). In this experiment section we use our EGNN to build an Equivariant Graph Autoencoder. We will explain how Graph Autoencoders can benefit from equivariance and we will show how our method outperforms standard GNN autoencoders in the provided datasets. This problem is particularly interesting since the embedding space can be scaled to larger dimensions and is not limited to a 3 dimensional Euclidean space.
The symmetry problem: The above stated autoencoder may seem straightforward to implement at first sight but in some cases there is a strong limitation regarding the symmetry of the graph. Graph Neural Networks are convolutions on the edges and nodes of a graph, i.e. the same function is applied to all edges and to all nodes. In some graphs (e.g. those defined only by its adjacency matrix) we may not have input features in the nodes, and for that reason the difference among nodes relies only on their edges or neighborhood topology. Therefore, if the neighborhood of two nodes is exactly the same, their encoded embeddings will be the same too. A clear example of this is a cycle graph (an example of a 4 nodes cycle graph is provided in Figure 3). When running a Graph Neural Network encoder on a node featureless cycle graph, we will obtain the exact same embedding for each of the nodes, which makes it impossible to reconstruct the edges of the original graph from the node embeddings. The cycle graph is a severe example where all nodes have the exact same neighborhood topology but these symmetries can be present in different ways for other graphs with different edge distributions or even when including node features if these are not unique.
Dataset: We generated community-small graphs (You et al., 2018; Liu et al., 2019) by running the original code from (You et al., 2018). These graphs contain nodes. We also generated a second dataset using the Erdos&Renyi generative model (Bollobás & Béla, 2001) sampling random graphs with an initial number of nodes and edge probability . We sampled graphs for training, for validation and for test for both datasets. Each graph is defined as and adjacency matrix .
Overfitting the training set: We explained the symmetry problem and we showed the EGNN outperforms other methods in the given datasets. Although we observed that adding noise to the GNN improves the results, it is difficult to exactly measure the impact of the symmetry limitation in these results independent from other factors such as generalization from the training to the test set. In this section we conduct an experiment where we train the different models in a subset of 100 Erdos&Renyi graphs and embedding size with the aim to overfit the data. We evaluate the methods on the training data. In this experiment the GNN is unable to fit the training data properly while the EGNN can achieve perfect reconstruction and Noise-GNN close to perfect. We sweep over different sparsity values from to since the symmetry limitation is more present in very sparse or very dense graphs. We report the F1 scores of this experiment in the right plot of Figure 4.
3 Molecular data — QM9
The QM9 dataset (Ramakrishnan et al., 2014) has become a standard in machine learning as a chemical property prediction task. The QM9 dataset consists of small molecules represented as a set of atoms (up to 29 atoms per molecule), each atom having a 3D position associated and a five dimensional one-hot node embedding that describe the atom type (H, C, N, O, F). The dataset labels are a variety of chemical properties for each of the molecules which are estimated through regression. These properties are invariant to translations, rotations and reflections on the atom positions. Therefore those models that are E(3) invariant are highly suitable for this task.
We imported the dataset partitions from (Anderson et al., 2019), 100K molecules for training, 18K for validation and 13K for testing. A variety of 12 chemical properties were estimated per molecule. We optimized and report the Mean Absolute Error between predictions and ground truth.
Conclusions
Acknowledgements
We would like to thank Patrick Forré for his support to formalize the invariance features identification proof.
References
Appendix A Equivariance Proof
Therefore, we have proven that rotating and translating results in the same rotation and translation on at the output of Equation 4.
Appendix B Re-formulation for velocity type inputs
In Appendix A we already proved the equivariance of our EGNN (Section 3) when not including vector type inputs. In its velocity type inputs variant we only replaced its coordinate updates (eq. 4) by Equation 7 that includes velocity. Since this is the only modification we will only prove that Equation 7 re-written below is equivariant.
First, we prove the first line preserves equivariance, that is we want to show:
Finally, it is straightforward to show the second equation is also equivariant, that is we want to show
Appendix C Implementation details
In this Appendix section we describe the implementation details of the experiments. First, we describe those parts of our model that are the same across all experiments. Our EGNN model from Section 3 contains the following three main learnable functions.
The edge function (eq. 3) is a two layers MLP with two Swish non-linearities: Input {LinearLayer() Swish() LinearLayer() Swish() } Output.
The coordinate function (eq. 4) consists of a two layers MLP with one non-linearity: {LinearLayer() Swish() LinearLayer() } Output
The node function (eq. 6) consists of a two layers MLP with one non-linearity and a residual connection:
[, ] {LinearLayer() Swish() LinearLayer() Addition() }
These functions are used in our EGNN across all experiments. Notice the GNN (eq. 2) also contains and edge operation and a node operation and respectively. We use the same functions described above for both the GNN and the EGNN such that comparisons are as fair as possible.
In the dynamical systems experiment we used a modification of the Charged Particle’s N-body (N=5) system from (Kipf et al., 2018). Similarly to (Fuchs et al., 2020), we extended it from 2 to 3 dimensions customizing the original code from (https://github.com/ethanfetaya/NRI) and we removed the virtual boxes that bound the particle’s positions. The sampled dataset consists of 3.000 training trajectories, 2.000 for validation and 2.000 for testing. Each trajectory has a duration of 1.000 timesteps. To move away from the transient phase, we actually generated trajectories of 5.000 time steps and sliced them from timestep to timestep (1.000 time steps into the future) such that the initial conditions are more realistic than the Gaussian Noise initialization from which they are initialized.
In our second experiment, we sweep from 100 to 50.000 training samples, for this we just created a new training partition following the same procedure as before but now generating 50.000 trajectories instead. The validation and test partition remain the same from last experiment.
All models are composed of 4 layers, the details for each model are the following.
EGNN: For the EGNN we use its variation that considers vector type inputs from Section 3.2. This variation adds the function to the model which is composed of two linear layers with one non-linearity: Input {LinearLayer() Swish() LinearLayer() } Output. Functions , and that define our EGNN are the same than for all experiments and are described at the beginning of this Appendix C.
GNN: The GNN is also composed of 4 layers, its learnable functions edge operation and node operation from Equation 2 are exactly the same as and from our EGNN introduced in Appendix C. We chose the same functions for both models to ensure a fair comparison. In the GNN case, the initial position and velocity from the particles is passed through a linear layer and inputted into the GNN first layer . The particle’s charges are inputted as edge attributes . The output of the GNN is passed through a two layers MLP that maps it to the estimated position.
Tensor Field Network: We used the Pytorch implementation from https://github.com/FabianFuchsML/se3-transformer-public. We swept over different hyper paramters, degree {2, 3, 4}, number of features {12, 24, 32, 64, 128}. We got the best performance in our dataset for degree 2 and number of features 32. We used the Relu activation layer instead of the Swish for this model since it provided better performance.
SE(3) Transformers: For the SE(3)-Transformer we used code from https://github.com/FabianFuchsML/se3-transformer-public. Notice this implementation has only been validated in the QM9 dataset but it is the only available implementation of this model. We swept over different hyperparamters degree {1, 2, 3, 4}, number of features 16, 32, 64 and divergence {1, 2}, along with the learning rate. We obtained the best performance for degree 3, number of features 64 and divergence 1. As in Tensor Field Networks we obtained better results by using the Relu activation layer instead of the Swish.
In Table 2 all models were trained for 10.000 epochs, batch size 100, Adam optimizer, the learning rate was fixed and independently chosen for each model. All models are 4 layers deep and the number of training samples was set to 3.000.
C.2 Implementation details for Graph Autoneoders
In this experiment we worked with Community Small (You et al., 2018) and Erdos&Renyi (Bollobás & Béla, 2001) generated datasets.
Community Small: We used the original code from (You et al., 2018) (https://github.com/JiaxuanYou/graph-generation) to generate a Community Small dataset. We sampled 5.000 training graphs, 500 for validation and 500 for testing.
Erdos&Renyi is one of the most famous graph generative algorithms. We used the ”gnp_random_graph(, )” function from (https://networkx.org/) that generates random graphs when povided with the number of nodes and the edge probability following the Erdos&Renyi model. Again we generated 5.000 graphs for training, 500 for validation and 500 for testing. We set the edge probability (or sparsity value) to and the number of nodes ranging from 7 to 16 deterministically uniformly distributed. Notice that edges are generated stochastically with probability , therefore, there is a chance that some nodes are left disconnected from the graph, ”gnp_random_graph(, )” function discards these disconnected nodes such that even if we generate graphs setting parameters to and the generated graphs may have less number of nodes.
Finally, in the graph autoencoding experiment we also overfitted in a small partition of 100 samples (Figure 4) for the Erdos&Renyi graphs described above. We reported results for different values ranging from to . For each value we generated a partition of 100 graphs with initial number of nodes between using the Erdos&Renyi generative model.
All experiments have been trained with learning rate , batch size 1, Adam optimizer, weight decay , 100 training epochs for the 5.000 samples sized datasets performing early stopping for the minimum Binary Cross Entropy loss in the validation partition. The overfitting experiments were trained for 10.000 epochs on the 100 samples subsets.
C.3 Implementation details for QM9
For QM9 (Ramakrishnan et al., 2014) we used the dataset partitions from (Anderson et al., 2019). We imported the dataloader from his code repository (https://github.com/risilab/cormorant) which includes his data-preprocessing. Additionally all properties have been normalized by substracting the mean and dividing by the Mean Absolute Deviation.
Our EGNN consists of 7 layers. Functions and are defined at the beginning of this Appendix C. Additionally, we use the module presented in Section 3.3 that infers the edges . This function is defined as a linear layer followed by a sigmoid: Input {Linear() sigmoid()} Output. Finally, the output of our EGNN is forwarded through a two layers MLP that acts node-wise, a sum pooling operation and another two layers MLP that maps the averaged embedding to the predicted property value, more formally: {Linear() Swish() Linear() Sum-Pooling() Linear() Swish() Linear} Property. The number of hidden features for all model hidden layers is 128.
We trained each property individually for a total of 1.000 epochs, we used Adam optimizer, batch size 96, weight decay , and cosine decay for the learning rate starting at at a lr= except for the Homo, Lumo and Gap properties where its initial value was set to .
Appendix D Further experiments
In this section we present an extension of the Graph Autoencoder experiment 5.2. In Table 4 we report the approximation error of the reconstructed graphs as the embedding dimensionality is reduced in the Community Small and Erdos&Renyi datasets for the GNN, Noise-GNN and EGNN models. For small embedding sizes () all methods perform poorly, but as the embedding size grows our EGNN significantly outperforms the others.
Appendix E Sometimes invariant features are all you need.
So without loss of generality, we may assume that . As a direct consequence . Now writing out the square:
And since , it follows that or equivalently written as dot product . Notice that this already shows that angles between pairs of points are the same.
At this moment, it might already be intuïtive that the collections of points are indeed identical. To finalize the proof formally we will construct a linear map for which we will show that (1) it maps every to and (2) that it is orthogonal. First note that from the angle equality it follows immediately that for every linear combination:
Let be the linear span of (so is the linear subspace of all linear combinations of ). Let be a basis of , where . Recall that one can define a linear map by choosing a basis, and then define for each basis vector where it maps to. Define a linear map from to by the transformation from the basis to for . Now pick any point and write it in its basis . We want to show or alternatively . Note that . Then:
Thus showing that for all , proving (1). Finally we want to show that is orthogonal, when restricted to . This follows since:
for the basis elements . This implies that is orthogonal (at least when restricted to ). Finally can be extended via an orthogonal complement of to the whole space. This concludes the proof for (2) and shows that is indeed orthogonal.