Directional Message Passing for Molecular Graphs

Johannes Gasteiger, Janek Groß, Stephan Günnemann

Introduction

In recent years scientists have started leveraging machine learning to reduce the computation time required for predicting molecular properties from a matter of hours and days to mere milliseconds. With the advent of graph neural networks (GNNs) this approach has recently experienced a small revolution, since they do not require any form of manual feature engineering and significantly outperform previous models (Gilmer et al. 2017; Schütt et al. 2017). GNNs model the complex interactions between atoms by embedding each atom in a high-dimensional space and updating these embeddings by passing messages between atoms. By predicting the potential energy these models effectively learn an empirical potential function. Classically, these functions have been modeled as the sum of four parts: (Leach 2001)

where EbondsE_{\text{bonds}} models the dependency on bond lengths, EangleE_{\text{angle}} on the angles between bonds, EtorsionE_{\text{torsion}} on bond rotations, i.e. the dihedral angle between two planes defined by pairs of bonds, and Enon-bondedE_{\text{non-bonded}} models interactions between unconnected atoms, e.g. via electrostatic or van der Waals interactions. The update messages in GNNs, however, only depend on the previous atom embeddings and the pairwise distances between atoms – not on directional information such as bond angles and rotations. Thus, GNNs lack the second and third terms of this equation and can only model them via complex higher-order interactions of messages. Extending GNNs to model them directly is not straightforward since GNNs solely rely on pairwise distances, which ensures their invariance to translation, rotation, and inversion of the molecule, which are important physical requirements.

Directional message passing, which allows GNNs to incorporate directional information by connecting recent advances in the fields of equivariance and graph neural networks as well as ideas from belief propagation and empirical potential functions such as Eq. 1.

Theoretically principled orthogonal basis representations based on spherical Bessel functions and spherical harmonics. Bessel functions achieve better performance than Gaussian radial basis functions while reducing the radial basis dimensionality by 4x or more.

The Directional Message Passing Neural Network (DimeNet): A novel GNN that leverages these innovations to set the new state of the art for molecular predictions and is suitable both for predicting molecular properties and for molecular dynamics simulations.

Related work

ML for molecules. The classical way of using machine learning for predicting molecular properties is combining an expressive, hand-crafted representation of the atomic neighborhood (Bartók et al. 2013) with Gaussian processes (Bartók et al. 2010; Bartók et al. 2017; Chmiela et al. 2017) or neural networks (Behler & Parrinello 2007). Recently, these methods have largely been superseded by graph neural networks, which do not require any hand-crafted features but learn representations solely based on the atom types and coordinates molecules (Duvenaud et al. 2015; Gilmer et al. 2017; Schütt et al. 2017; Hy et al. 2018; Unke & Meuwly 2019). Our proposed message embeddings can also be interpreted as directed edge embeddings or embeddings on the line graph (Chen et al. 2019b). (Undirected) edge embeddings have already been used in previous GNNs for molecules (Jørgensen et al. 2018; Chen et al. 2019a). However, these GNNs use both node and edge embeddings and do not leverage any directional information.

Graph neural networks. GNNs were first proposed in the 90s (Baskin et al. 1997; Sperduti & Starita 1997) and 00s (Gori et al. 2005; Scarselli et al. 2009). General GNNs have been largely inspired by their application to molecular graphs and have started to achieve breakthrough performance in various tasks at around the same time the molecular variants did (Kipf & Welling 2017; Gasteiger et al. 2019; Zambaldi et al. 2019). Some recent progress has been focused on GNNs that are more powerful than the 1-Weisfeiler-Lehman test of isomorphism (Morris et al. 2019; Maron et al. 2019). However, for molecular predictions these models are significantly outperformed by GNNs focused on molecules (see Sec. 7). Some recent GNNs have incorporated directional information by considering the change in local coordinate systems per atom (Ingraham et al. 2019). However, this approach breaks permutation invariance and is therefore only applicable to chain-like molecules (e.g. proteins).

Equivariant neural networks. Group equivariance as a principle of modern machine learning was first proposed by Cohen & Welling 2016. Following work has generalized this principle to spheres (Cohen et al. 2018), molecules (Thomas et al. 2018), volumetric data (Weiler et al. 2018), and general manifolds (Cohen et al. 2019). Equivariance with respect to continuous rotations has been achieved so far by switching back and forth between Fourier and coordinate space in each layer (Cohen et al. 2018) or by using a fully Fourier space model (Kondor et al. 2018; Anderson et al. 2019). The former introduces major computational overhead and the latter imposes significant constraints on model construction, such as the inability of using non-linearities. Our proposed solution does not suffer from either of those limitations.

Requirements for molecular predictions

Symmetries and invariances. All molecular predictions must obey some basic laws of physics, either explicitly or implicitly. One important example of such are the fundamental symmetries of physics and their associated invariances. In principle, these invariances can be learned by any neural network via corresponding weight matrix symmetries (Ravanbakhsh et al. 2017). However, not explicitly incorporating them into the model introduces duplicate weights and increases training time and complexity. The most essential symmetries are translational and rotational invariance (follows from homogeneity and isotropy), permutation invariance (follows from the indistinguishability of particles), and symmetry under parity, i.e. under sign flips of single spatial coordinates.

Molecular dynamics. Additional requirements arise when the model should be suitable for molecular dynamics (MD) simulations and predict the forces Fi{\bm{F}}_{i} acting on each atom. The force field is a conservative vector field since it must satisfy conservation of energy (the necessity of which follows from homogeneity of time (Noether 1918)). The easiest way of defining a conservative vector field is via the gradient of a potential function. We can leverage this fact by predicting a potential instead of the forces and then obtaining the forces via backpropagation to the atom coordinates, i.e. Fi(X,z)=−∂∂xifθ(X,z){\bm{F}}_{i}({\bm{X}},{\bm{z}})=-\frac{\partial}{\partial{\bm{x}}_{i}}f_{\theta}({\bm{X}},{\bm{z}}). We can even directly incorporate the forces in the training loss and directly train a model for MD simulations (Pukrittayakamee et al. 2009):

where the target t^=E^\hat{t}=\hat{E} is the ground-truth energy (usually available as well), F^\hat{{\bm{F}}} are the ground-truth forces, and the hyperparameter ρ\rho sets the forces’ loss weight. For stable simulations Fi{\bm{F}}_{i} must be continuously differentiable and the model fθf_{\theta} itself therefore twice continuously differentiable. We hence cannot use discontinuous transformations such as ReLU non-linearities. Furthermore, since the atom positions X{\bm{X}} can change arbitrarily we cannot use pre-computed auxiliary information Θ{\bm{\Theta}} such as bond types.

Directional message passing

with the update function fupdatef_{\text{update}} and the interaction function fintf_{\text{int}}, which are both commonly implemented using neural networks. The edge embeddings e(ij)(l){\bm{e}}_{(ij)}^{(l)} usually only depend on the interatomic distances, but can also incorporate additional bond information (Gilmer et al. 2017) or be recursively updated in each layer using the neighboring atom embeddings (Jørgensen et al. 2018).

Directional embeddings. We solve this problem by noting that an atom by itself is rotationally invariant. This invariance is only broken by neighboring atoms that interact with it, i.e. those inside the cutoff cc. Since each neighbor breaks up to one rotational invariance they also introduce additional degrees of freedom, which we need to represent in our model. We can do so by generating a separate embedding mji{\bm{m}}_{ji} for each atom ii and neighbor jj by applying the same learned filter in the direction of each neighboring atom (in contrast to equivariant CNNs, which apply filters in fixed, global directions). These directional embeddings are equivariant with respect to global rotations since the associated directions rotate with the molecule and hence conserve the relative directional information between neighbors.

Message embeddings. The directional embedding mji{\bm{m}}_{ji} associated with the atom pair jiji can be thought of as a message being sent from atom jj to atom ii. Hence, in analogy to belief propagation, we embed each atom ii using a set of incoming messages mji{\bm{m}}_{ji}, i.e. hi=∑j∈Nimji{\bm{h}}_{i}=\sum_{j\in\mathcal{N}_{i}}{\bm{m}}_{ji}, and update the message mji{\bm{m}}_{ji} based on the incoming messages mkj{\bm{m}}_{kj} (Yedidia et al. 2003). Hence, as illustrated in Fig. 1, we define the update function and aggregation scheme for message embeddings as

where eRBF(ji){\bm{e}}_{\text{RBF}}^{(ji)} denotes the radial basis function representation of the interatomic distance djid_{ji}, which will be discussed in Sec. 5. We found this aggregation scheme to not only have a nice analogy to belief propagation, but also to empirically perform better than alternatives. Note that since fintf_{\text{int}} now incorporates the angle between atom pairs, or bonds, we have enabled our model to directly learn the angular potential EangleE_{\text{angle}}, the second term in Eq. 1. Moreover, the message embeddings are essentially embeddings of atom pairs, as used by the provably more powerful GNNs based on higher-order Weisfeiler-Lehman tests of isomorphism. Our model can therefore provably distinguish molecules that a regular GNN cannot (e.g. the previous example of a hexagonal and two triangular molecules) (Morris et al. 2019).

Physically based representations

Representing distances and angles. For the interaction function fintf_{\text{int}} in Eq. 4 we use a joint representation aSBF(kj,ji){\bm{a}}_{\text{SBF}}^{(kj,ji)} of the angles α(kj,ji)\alpha_{(kj,ji)} between message embeddings and the interatomic distances dkj=∥xk−xj∥2d_{kj}=\|{\bm{x}}_{k}-{\bm{x}}_{j}\|_{2}, as well as a representation eRBF(ji){\bm{e}}_{\text{RBF}}^{(ji)} of the distances djid_{ji}. Earlier works have used a set of Gaussian radial basis functions to represent interatomic distances, with tightly spaced means that are distributed e.g. uniformly (Schütt et al. 2017) or exponentially (Unke & Meuwly 2019). Similar in spirit to the functional bases used by steerable CNNs (Cohen & Welling 2017; Cheng et al. 2019) we propose to use an orthogonal basis instead, which reduces redundancy and thus improves parameter efficiency. Furthermore, a basis chosen according to the properties of the modeled system can even provide a helpful inductive bias. We therefore derive a proper basis representation for quantum systems next.

From Schrödinger to Fourier-Bessel. To construct a basis representation in a principled manner we first consider the space of possible solutions. Our model aims at approximating results of density functional theory (DFT) calculations, i.e. results given by an electron density <Ψ(d)∣Ψ(d)>\left<\Psi({\bm{d}})|\Psi({\bm{d}})\right>, with the electron wave function Ψ(d)\Psi({\bm{d}}) and d=xk−xj{\bm{d}}={\bm{x}}_{k}-{\bm{x}}_{j}. The solution space of Ψ(d)\Psi({\bm{d}}) is defined by the time-independent Schrödinger equation (−ℏ22m∇2+V(d))Ψ(d)=EΨ(d)\left(-\frac{\hbar^{2}}{2m}\nabla^{2}+V({\bm{d}})\right)\Psi({\bm{d}})=E\Psi({\bm{d}}), with constant mass mm and energy EE. We do not know the potential V(d)V({\bm{d}}) and so choose it in an uninformative way by simply setting it to 0 inside the cutoff distance cc (up to which we pass messages between atoms) and to ∞\infty outside. Hence, we arrive at the Helmholtz equation (∇2+k2)Ψ(d)=0(\nabla^{2}+k^{2})\Psi({\bm{d}})=0, with the wave number k=2mEℏk=\frac{\sqrt{2mE}}{\hbar} and the boundary condition Ψ(c)=0\Psi(c)=0 at the cutoff cc. Separation of variables in polar coordinates (d,α,φ)(d,\alpha,\varphi) yields the solution (Griffiths & Schroeter 2018)

with n∈[1..NRBF]n\in[1\mathinner{\ldotp\ldotp}N_{\text{RBF}}]. Both of these bases are purely real-valued and orthogonal in the domain of interest. They furthermore enable us to bound the highest-frequency components by ωα≤NSHBF2π\omega_{\alpha}\leq\frac{N_{\text{SHBF}}}{2\pi}, ωdkj≤NSRBFc\omega_{d_{kj}}\leq\frac{N_{\text{SRBF}}}{c}, and ωdji≤NRBFc\omega_{d_{ji}}\leq\frac{N_{\text{RBF}}}{c}. This restriction is an effective way of regularizing the model and ensures that predictions are stable to small perturbations. We found NSRBF=6N_{\text{SRBF}}=6 and NRBF=16N_{\text{RBF}}=16 radial basis functions to be more than sufficient. Note that NRBFN_{\text{RBF}} is 4x lower than PhysNet’s 64 (Unke & Meuwly 2019) and 20x lower than SchNet’s 300 radial basis functions (Schütt et al. 2017).

Directional Message Passing Neural Network (DimeNet)

The Directional Message Passing Neural Network’s (DimeNet) design is based on a streamlined version of the PhysNet architecture (Unke & Meuwly 2019), in which we have integrated directional message passing and spherical Fourier-Bessel representations. DimeNet generates predictions that are invariant to atom permutations and translation, rotation and inversion of the molecule. DimeNet is suitable both for the prediction of various molecular properties and for molecular dynamics (MD) simulations. It is twice continuously differentiable and able to learn and predict atomic forces via backpropagation, as described in Sec. 3. The predicted forces fulfill energy conservation by construction and are equivariant with respect to permutation and rotation. Model differentiability in combination with basis representations that have bounded maximum frequencies furthermore guarantees smooth predictions that are stable to small deformations. Fig. 4 gives an overview of the architecture.

where ∥\| denotes concatenation and the weight matrix W{\bm{W}} and bias b{\bm{b}} are learnable.

Interaction block. The embedding block is followed by multiple stacked interaction blocks. This block implements fintf_{\text{int}} and fupdatef_{\text{update}} of Eq. 4 as shown in Fig. 4. Note that the 2D representation aSBF(kj,ji){\bm{a}}_{\text{SBF}}^{(kj,ji)} is first transformed into an NbilinearN_{\text{bilinear}}-dimensional representation via a linear layer. The main purpose of this is to make the dimensionality of aSBF(kj,ji){\bm{a}}_{\text{SBF}}^{(kj,ji)} independent of the subsequent bilinear layer, which uses a comparatively large Nbilinear×F×FN_{\text{bilinear}}\times F\times F-dimensional weight tensor. We have also experimented with using a bilinear layer for the radial basis representation, but found that the element-wise multiplication eRBF(ji)W⊙mkj{\bm{e}}_{\text{RBF}}^{(ji)}{\bm{W}}\odot{\bm{m}}_{kj} performs better, which suggests that the 2D representations require more complex transformations than radial information alone. The interaction block transforms each message embedding mji{\bm{m}}_{ji} using multiple residual blocks, which are inspired by ResNet (He et al. 2016) and consist of two stacked dense layers and a skip connection.

Output block. The message embeddings after each block (including the embedding block) are passed to an output block. The output block transforms each message embedding mji{\bm{m}}_{ji} using the radial basis eRBF(ji){\bm{e}}_{\text{RBF}}^{(ji)}, which ensures continuous differentiability and slightly improves performance. Afterwards the incoming messages are summed up per atom ii to obtain hi=∑jmji{\bm{h}}_{i}=\sum_{j}{\bm{m}}_{ji}, which is then transformed using multiple dense layers to generate the atom-wise output ti(l)t_{i}^{(l)}. These outputs are then summed up to obtain the final prediction t=∑i∑lti(l)t=\sum_{i}\sum_{l}t_{i}^{(l)}.

Experiments

Models. For hyperparameter choices and training setup see Appendix B. We use 6 state-of-the-art models for comparison: SchNet (Schütt et al. 2017), PhysNet (results based on the reference implementation) (Unke & Meuwly 2019), provably powerful graph networks (PPGN, results provided by the original authors) (Maron et al. 2019), MEGNet-simple (without auxiliary information) (Chen et al. 2019a), Cormorant (Anderson et al. 2019), and symmetrized gradient-domain machine learning (sGDML) (Chmiela et al. 2018). Note that sGDML cannot be used for QM9 since it can only be trained on a single molecule.

Conclusion

In this work we have introduced directional message passing, a more powerful and expressive interaction scheme for molecular predictions. Directional message passing enables graph neural networks to leverage directional information in addition to the interatomic distances that are used by normal GNNs. We have shown that interatomic distances can be represented in a principled and effective manner using spherical Bessel functions. We have furthermore shown that this representation can be extended to directional information by leveraging 2D spherical Fourier-Bessel basis functions. We have leveraged these innovations to construct DimeNet, a GNN suitable both for predicting molecular properties and for use in molecular dynamics simulations. We have demonstrated DimeNet’s performance on QM9 and MD17 and shown that our contributions are the essential ingredients that enable DimeNet’s state-of-the-art performance. DimeNet directly models the first two terms in Eq. 1, which are known as the important “hard” degrees of freedom in molecules (Leach 2001). Future work should aim at also incorporating the third and fourth terms of this equation. This could improve predictions even further and enable the application to molecules much larger than those used in common benchmarks like QM9.

This research was supported by the German Federal Ministry of Education and Research (BMBF), grant no. 01IS18036B, and by the Deutsche Forschungsgemeinschaft (DFG) through the Emmy Noether grant GU 1409/2-1 and the TUM International Graduate School of Science and Engineering (IGSSE), GSC 81. The authors of this work take full responsibilities for its content.

References

Appendix A Indistinguishable molecules

Appendix B Experimental setup

The model architecture and hyperparameters were optimized using the QM9 validation set. We use 6 stacked interaction blocks and embeddings of size F=128F=128 throughout the model. For the basis functions we choose NSHBF=7N_{\text{SHBF}}=7 and NSRBF=NRBF=6N_{\text{SRBF}}=N_{\text{RBF}}=6. For the weight tensor in the interaction block we use Nbilinear=8N_{\text{bilinear}}=8. We did not find the model to be very sensitive to these values as long as they were large enough (i.e. at least 4).

Appendix C Summary statistics

We summarize the results across different targets using the mean standardized MAE

with target index mm, number of targets M=12M=12, dataset size NN, ground truth values t^(m)\hat{{\bm{t}}}^{(m)}, model fθ(m)f_{\theta}^{(m)}, inputs Xi{\bm{X}}_{i} and zi{\bm{z}}_{i}, and standard deviation σm\sigma_{m} of t^(m)\hat{{\bm{t}}}^{(m)}. Std. MAE reflects the average error compared to the standard deviation of each target. Since this error is dominated by a few difficult targets (e.g. ϵHOMO\epsilon_{\text{HOMO}}) we also report logMAE, which reflects every relative improvement equally but is sensitive to outliers, such as SchNet’s result on <R2>\left<R^{2}\right>.

Appendix D DimeNet filters

To illustrate the filters learned by DimeNet we separate the spatial dependency in the interaction function fintf_{\text{int}} via

where WRBF{\bm{W}}_{\text{RBF}}, WSBF{\bm{W}}_{\text{SBF}}, and W{\bm{\mathsfit{W}}} are learned weight matrices/tensors, eRBF(d){\bm{e}}_{\text{RBF}}(d) is the radial basis representation, and aSBF(d,α){\bm{a}}_{\text{SBF}}(d,\alpha) is the 2D spherical Fourier-Bessel representation. Fig. 7 shows how the first 15 elements of ffilter2,n(d,α)f_{\text{filter2},n}(d,\alpha) vary with dd and α\alpha when choosing the tensor slice n=1n=1 (with α=0\alpha=0 at the top of the figure).

Appendix E Multi-target results