Fast and Uncertainty-Aware Directional Message Passing for Non-Equilibrium Molecules

Johannes Gasteiger, Shankari Giri, Johannes T. Margraf, Stephan Günnemann

Introduction

Modern machine learning models for molecular property prediction typically focus on molecules in equilibrium (e.g. QM9 ) or close to the equilibrium (e.g. MD17 , ANI-1 , QM7-X ). However, this precludes their application to the dynamics during chemical reactions, which involve transition states far away from the equilibrium. Making reliable predictions for these states requires models that are able to cover a much broader range of chemical and configurational space, i.e. including open-shell electronic structures, stretched bonds and distorted angles. In this work we aim at making progress on this problem from three directions.

Second, we develop a new dataset that contains highly reactive non-equilibrium systems. The new COLL dataset contains 140 000140\,000 configurations of pairs of molecules reacting at high kinetic energies. It only consists of small molecules but covers the space of reactions much better and includes a significantly wider range of energies and forces than previous benchmarks, as shown in Fig. 4.

Due to the vast number of possible non-equilibrium configurations it is crucial that we are able to detect when we move out of the region covered by the training data and react appropriately (e.g. via active learning). To achieve this we investigate ensembling and mean-variance estimation . We conclude that both are insufficient, due to their overhead and inability of reliably predicting the energy and force uncertainties.

DimeNet++

DimeNet. DimeNet is a recently proposed Graph Neural Network (GNN) for molecular property prediction . It improves upon regular GNNs in two ways. Normal GNNs represent each atom ii separately via its embedding hi{\bm{h}}_{i} and update these in each layer ll via message passing. DimeNet instead embeds and updates the messages between atoms mji{\bm{m}}_{ji}, which enables it to consider directional information (via bond angles α(kj,ji)\alpha_{(kj,ji)}) as well as interatomic distances djid_{ji}. DimeNet furthermore embeds distances and angles jointly using a spherical 2D Fourier-Bessel basis, resulting in the update

where fupdatef_{\text{update}} denotes the update function, fintf_{\text{int}} the interaction function, eRBF(ji){\bm{e}}_{\text{RBF}}^{(ji)} the radial basis function (RBF) representation of djid_{ji} and aSBF(kj,ji){\bm{a}}_{\text{SBF}}^{(kj,ji)} the spherical basis function (SBF) representation of dkjd_{kj} and α(kj,ji)\alpha_{(kj,ji)}. In this work we do not touch either of those contributions and instead focus on the model architecture. The updated DimeNet++ architecture is illustrated in Fig. 1.

Fast interactions. We therefore first focus on the expensive “directional message passing” block. It is DimeNet’s centerpiece, modelling the interaction between embeddings mkj{\bm{m}}_{kj} and basis representations eRBF(ji){\bm{e}}_{\text{RBF}}^{(ji)} and aSBF(kj,ji){\bm{a}}_{\text{SBF}}^{(kj,ji)}. As such, it requires an adequately expressive transformation. The original DimeNet accomplishes this with a bilinear layer, as shown in Fig. 2. Unfortunately, this layer is very expensive, which is exacerbated by being used in the model’s most costly component. We alleviate this by replacing it with a simple Hadamard product and compensate for the loss in expressiveness by adding multilayer perceptrons (MLPs) for the basis representations. This recovers the original accuracy at a fraction of the computational cost (see Section 5).

Embedding hierarchy. We can directly leverage the fact that certain parts of the model use a higher number of embeddings by reducing the embedding size in these parts via down- and upprojection layers W↓{\bm{W}}_{\downarrow} and W↑{\bm{W}}_{\uparrow}. This both accelerates the model and removes information bottlenecks, since we no longer aggregate information to a smaller number of equally sized embeddings.

Other improvements. We furthermore found that using 4 layers performs en par with the original 6 for U0U_{0}. Moreover, larger batch sizes significantly slowed down convergence, and mixed precision caused the model’s precision to break down completely. Considering that DimeNet’s relative error is below float16’s machine precision (5⋅10−45\cdot 10^{-4}), the latter might be expected.

COLL Dataset

The COLL dataset consists of configurations taken from molecular dynamics simulations of molecular collisions. To this end, collision simulations were performed with the cost-effective semiempirical GFN2-xTB method . Subsequently, energies and forces for 140 000140\,000 random snapshots taken from these trajectories were recomputed with density functional theory (DFT). These calculations were performed with the revPBE functional and def2-TZVP basis, including D3 dispersion corrections .

Exemplary structures from the COLL set are shown in Fig. 4. Unlike established molecular benchmark sets (e.g. QM9), which consist of equilibrium or near-equilibrium configurations, the structures in COLL can be highly distorted. In particular, stretched bonds and angles, as well as open-shell electronic structures are prevalent. All calculations are preformed with broken spin-symmetry.

Uncertainty Quantification

The vast number of non-equilibrium states reachable in high-energy molecular dynamics simulations (such as reactions) means that systems will often move outside the space covered by our training set. We therefore need a reliable way of detecting a degradation in predictive performance. Most uncertainty quantification (UQ) methods are focused on providing an uncertainty estimate for the direct prediction . However, out-of-equilibrium dynamics require uncertainty estimates for both the energy EE and the force F=−∂E∂x{\bm{F}}=-\frac{\partial E}{\partial{\bm{x}}}. Many non-differentiable methods (e.g. combining GNNs with a random forest) are therefore not applicable. Ensembling is a notable exception but introduces a large computational overhead since we need to calculate predictions using multiple separate models.

Even if the method is differentiable it might only provide a mean and standard deviation, i.e. μE\mu_{E} and σE\sigma_{E} (e.g. mean-variance estimation (MVE)). This allows us to obtain the force prediction via

However, performing the same operation on σE\sigma_{E} does not yield the analogous result:

There is thus no general way of estimating σF\sigma_{\bm{F}} for these kinds of models. Instead, we have to rely on σE\sigma_{E} as the uncertainty measure and hope that it correlates with the force error.

Experiments

DimeNet++ In Table 1 we evaluate each of the proposed DimeNet improvements separately on the U0U_{0} validation set of QM9 . We see that each change either reduces the runtime or improves the error. Exchanging the bilinear layer for a Hadamard product has by far the largest impact, single-handedly decreasing the runtime by a factor of 5. Interestingly, decreasing the embedding size both accelerates the model and improves the accuracy. This is either due to the additional down- and upprojection layers improving expressiveness or to the smaller embedding size improving generalization.

We evaluate the final DimeNet++ model on all QM9 targets and compare it to the state-of-the-art models SchNet , MGCN , and DeepMoleNet . Table 2 shows that DimeNet++ performs better for most targets and best overall, in addition to being 8x faster than DimeNet. OrbNet performs better on U0U_{0}, but has not published results for the other properties. Note that the DFT-based representations introduced by DeepMoleNet and OrbNet can also be incorporated into DimeNet++.

Uncertainty quantification. Ensembling and MVE both struggle with estimating the energy uncertainty, as shown for DimeNet++ in Table 3. The force error is very well estimated by the ensemble, but not by the energy uncertainty – especially for MVE. The energy uncertainty is thus not as good a proxy for the force error as one would expect. While the ensemble does perform decently, its computational overhead is still considerable. Reliable and fast uncertainty estimates thus remain an important direction for future work.

Acknowledgments and Disclosure of Funding

This research was supported by the TUM International Graduate School of Science and Engineering (IGSSE), GSC 81. The authors of this work take full responsibilities for its content.

References