Reliable Graph Neural Networks for Drug Discovery Under Distributional Shift
Kehang Han, Balaji Lakshminarayanan, Jeremiah Liu
Introduction
In recent years, Graph Neural Networks (GNNs) have illustrated remarkable performance in scientific applications such as physics (Battaglia et al., 2016; Sanchez-Gonzalez et al., 2018), biomedical science (Fout et al., 2017), and computational chemistry (Duvenaud et al., 2015; Kearnes et al., 2016). One important example is the early-phase drug discovery, where GNNs have shown promise in critical tasks such as hit-finding and liability screening (i.e., predicting the binding affinity and the toxicity of candidate drug molecules, respectively) (McCloskey et al., 2020; Siramshetty et al., 2020).
However, a key reliability concern that hinders the adoption of GNN in real practice is the overconfident mispredictions. For example, in liability screening, an Overconfident False Negative (OFN hereafter) prediction leads the GNN model to mark a toxic molecule as safe, causing severe consequences by leaking it to the next stage of drug development. Such concern is further exacerbated by a second key characteristic of the drug discovery tasks: data distributional shift; drug discovery tasks often explicitly evaluate novel molecules by moving into regions in the feature space that are not previously represented in training data. As a result, the testing molecules are characteristically distinct from the training data, and can carry novel toxic signatures that were previously unseen by the model. A model with just high in-distribution accuracy is not sufficient under those situations; quantifying model reliability and designing techniques to improve robustness against overconfidence under distributional shift become especially relevant.
Currently, there lacks a drug discovery benchmark that targets realistic concerns about model reliability under distributional shifts. Thus, we introduce CardioTox, a data benchmark based on a real-world drug discovery problem and is compiled from 9K+ drug-like molecules from ChEMBL, NCATS and FDA validation databases (Siramshetty et al., 2020). To evaluate model reliability, we generate additional molecular annotations and propose novel metrics to measure the models against the real-world standards around the responsible application of GNN. This is our first contribution.
Using CardioTox, our second contribution is an exploratory study on the root causes behind the overconfident mispredictions under distributional shift, and principled modeling approaches to mitigate it. In particular, we observe that many overconfidently mispredicted molecules are structurally distinct from training data (i.e., they are "far" from training data based on molecule-fingerprint graph distance (Rogers and Hahn, 2010), see Section 3). This failure mode suggests that improving the distance-awareness of a GNN model would be an effective solution: a test molecule that is far away from the decision boundary (e.g., a toxic but novel molecule, due to its novelty, may possess few prior toxic signatures defining the decision boundary) should still get low confidence if it’s distant from training data. To this end, the recently proposed Spectral-normalized Neural Gaussian Processes (SNGP) (Liu et al., 2020) demonstrates a concrete approach to achieve this goal. Specifically, it imposes a distance-preserving regularization (i.e., spectral normalization) to the feature extractor, and replaces the dense output layer with a distance-aware classifier (i.e., random-feature Gaussian process). SNGP has shown promising robustness improvements in vision and language problems. To our best knowledge, this is the first study to bring the distance-aware design principle to GNN models to to improve reliability performance on the molecule graph.
To summarize, our contributions are the following (Appendix E summarizes the related work):
Data: We introduce CardioTox, a real-world drug discovery dataset with multiple test sets (IID and distribution-shifted ones), to the robustness community to facilitate reliability research of graph models. The distributional shift challenge reflected in CardioTox is realistically faced by the field: test domain often comes from a data source that’s different from the training source and has considerable amount of novel molecule graph structures (see also Figure S3 in Appendix D). We further design distance-based data splits for CardioTox to quantitatively measure model’s distance-awareness.
Model: We develop an end-to-end trainable GNN-SNGP architecture as well as its ablated version GNN-GP. Empirically, this method outperforms its base architecture in not only accuracy (e.g., AUROC) but also robustness (e.g., ECE) especially in reducing overconfident mispredictions. Measured by CardioTox’s distance-based data splits, GNN-SNGP shows higher distance-awareness under data shifts than the baselines, which explains its ability in overconfident misprediction reduction. Our implementation together with CardioTox dataset will be open sourced via Github.https://github.com/google/uncertainty-baselines/tree/main/baselines/drug_cardiotoxicity
Ablation Study: We carry out extensive ablation studies which confirm that the robustness improvements come from both distance-aware classifier (GP-layer) and the distance-preserving latent representations. We also investigate the generalizability of this approach by testing on other graph modeling domain: molHIV, BBBP and BACE and obtained consistent results.
Methods
We use a vanilla Message Passing Neural Network (MPNN) (Gilmer et al., 2017) as the GNN baseline in this study. Specifically, the message function is modeled by a dense layer , where are node features for node respectively, is the edge feature vector between nodes , and is the weight matrix. After getting aggregated message for each node, the hidden node feature is updated via a Gated Recurrent Unit . Finally, we read out graph level representation by , where , are two weight matrices and is the sigmoid function. The graph level representation is fed into a final dense layer to generate logits.
GNN-GP: improving distance-awareness of classifier
Following Liu et al. (2020) who proposed Gaussian process layer (GP-layer hereafter) in vision and language models, we introduce it to the graph domain and developed GNN-GP model to increase distance-awareness. Figure S1 (Appendix A) shows the high-level architectural changes.
As Figure 1 shows in detail, Gaussian processes are made end-to-end trainable with GNN by approximating GP kernel function via random Fourier feature generation (Rahimi and Recht, 2007). Specifically, coefficients are randomly sampled at initialization and kept fixed during training. The distribution that are sampled from depends on what Gaussian processes kernel we’d like to approximate. Take the Gaussian kernel as an example, we have . During inference time, each sample would get logit predictions as well logit variances, both of which are utilized to compute predictive probabilities via mean-field approximation (Lu et al., 2020).
GNN-SNGP: incorporating distance-preserving feature extraction
Due to feature collapse in feature extraction (Liu et al., 2020; van Amersfoort et al., 2021), neural representation may not faithfully preserve distance in the input manifold. Liu et al. (2020) propose to preserve input distance in feature extraction by applying Spectral Normalization (SN) (Gouk et al., 2021; Miyato et al., 2018) to the residual networks. As a result, combining SN and GP would increase the model’s overall distance-awareness, helping OOD detection and other robustness metrics.
Since our vanilla MPNN model (GNN baseline) does not have residual connections, we create GNN-SNGP through two following changes (Figure 1). First, we model the message function via a dense layer with SN: where is the dense layer weight matrix whose spectral norm is regularized. Second, we add residual connection to the message passing layer through the node update function: .
CardioTox: Drug cardiotoxicity under distributional shift
In early drug discovery stages, two types of models are widely used: hit-finding model and anti-target (i.e., liability) model. The former scores a molecule based on its potential of making a drug (e.g., whether a molecule could bind tightly to the disease-causing protein target), while the latter aims to filter out those molecules with potentially high liability (e.g., whether a molecule could be toxic). One such liability is cardiotoxicity, which occurs when a molecule inhibits hERG, a protein target that is related to heart-rhythm control (Smith et al., 1996; Vandenberg et al., 2012). Due to previous failing incidents, FDA has required new drugs to pass drug cardiotoxicity examination (Center for Drug Evaluation and Research, 2005).
Thus in this work, we introduce CardioTox, a benchmark based on a real-world drug discovery problem and is compiled from 9K+ drug-like molecules from ChEMBL and NCATS databases (Siramshetty et al., 2020). To evaluate GNN model reliability, we add graph structural information (e.g., node features) and set up three test sets that reflect distributional shifts: Test-IID is a set that’s sampled from the same distribution as the Train set. Test-OOD1 and Test-OOD2 are molecule sets coming from NCATS and FDA respectively: 84% of Test-OOD1 and 82% of Test-OOD2 are novel molecules that are distant from Train set (Figure S3). We further generate additional molecular annotations: close sample or far sample based on Tanimoto fingerprint distance (Bajusz et al., 2015) to the train set (Appendix D) . This setup allows us to quantitatively assess models’ distance-awareness.
Accuracy and robustness performance Table 1 shows that GNN-GP outperforms GNN baseline in both AUROC and robustness metrics for the CardioTox task. This is the case for both the in-distribution (Test-IID) and the shifted test sets (Test-OOD1 and Test-OOD2). GNN-SNGP shows additional gains in robustness. If resource allows, Deep Ensemble (Lakshminarayanan et al., 2017) of GNN-SNGP further boosts performance, eliminating all OFNs in Test-OOD2. It is worth noting that our GNN-SNGP ensemble has outperformed previous state-of-the-art neural models (Siramshetty et al., 2020) in AUROC performance.
Why does distance-awareness help reduce overconfident mispredictions? We observe a significant portion of OFNs are distant from the train set (i.e., 60% of them have Tanimoto distance (Bajusz et al., 2015), also see Figure 2(b)). This could happen when a GNN model lacks of distance-awareness: a toxic but novel molecule can be far away from the wrong side of model’s decision boundary, due to lacking known toxic signatures that is present in train set. Without being aware of the distance to train set, the GNN baseline model tends to base its prediction on distance to the decision boundary and give high confidence for such cases.
To this end, GNN-GP leverages GP’s distance-awareness and is able to naturally incorporate distance into predictive uncertainty. As shown in Table 1, the SN version GNN-SNGP achieves highest distance-awareness (DA-AUC) in the two data-source-shifted test sets, correlating well with OFNs reduction. Figure 2(b) shows a decreasing trend (from GNN baseline to GNN-GP to GNN-SNGP) of the percentage of distant samples among OFNs. Overall, using GP is able to improve uncertainty estimate for over 80% of the baseline OFNs while also improving calibration (Figure S4).
Additional ablations to assess relative contributions of SN and GP In order to understand relative performance contributions of the two modeling components (i.e., distance-preserving feature extractor vs distance-aware classifier), we carry out an extensive ablation study in Appendix H. We take the latent representations learned by GNN baseline, GNN-GP and GNN-SNGP (increasing distance-preservation) and feed them to classifiers with increasing distance-awareness: Dense layer, GP-layer, exact Gaussian processes classifier (GPC hereafter). Table S2 suggests there’s synergy between distance-preservation of neural representation and distance-awareness of the classifier. With the least distance-aware classifier (i.e., Dense layer), AUROC drops when increasing distance-preservation in neural representation. With the most distance-aware classifier (i.e., exact GPC), increasing neural representation’s distance-preservation benefits accuracy, robustness as well as overconfidence reduction. We find a GNN model achieves best performances when both are present in the architecture (GNN-SNGP embeddings with GPC). Interesting, naively increasing distance-preservation along is not sufficient in guaranteeing good generalization; we observe that the pre-defined representation (i.e., the molecule fingerprint FP) gives relatively low AUROC under data shifts (e.g., on Test-OOD2) despite being perfectly distance-preserving. This is related to the trade-off in representation learning between dimension reduction and information preservation. Appendix H discusses in further detail.
Additional results on existing benchmarks Consistent with the results obtained in CardioTox, GNN-GP also outperforms GNN baseline on three established graph classification benchmarks (Wu et al., 2018): molHIV, BBBP and BACE (see Appendix G). We can make a few observations on the results in Table S1. First, GNN-GP consistently achieves higher AUROC, improves calibration as well as reduces overconfident mispredictions across all the benchmarks than the baseline. Second, with limited accuracy performance drop, the spectral normalized version GNN-SNGP can further improve model robustness compared with GNN-GP.
Conclusion
In this study, we introduce GNN-GP and GNN-SNGP together with a new benchmark CardioTox from drug discovery setting. Through evaluation on four datasets, we demonstrate their effectiveness in reducing overconfident mispredictions and making better calibrated GNN models without sacrificing accuracy performance. The improved robustness appears to come from the boost in distance-awareness. We further discover that the embedding space induced by SN and GP addition improves distance-preservation over its base architecture and is one major factor to bring improvements in accuracy, general robustness and overconfidence performance.
Moving forward, it would be interesting to carry out a broader empirical study with other base GNN architectures such as GAT (Veličković et al., 2017) and PNA (Corso et al., 2020). Another interesting direction, which is already initiated by the ablation study detailed in Appendix H, is to find a quantitative way to measure distance-preservation within learned representation and understand the trade-off between the preservation and task-specific compression via the lens of accuracy and robustness performance.
References
Appendix A GNN-SNGP architecture
We present high-level architectural changes based on GNN baseline model in Figure S1: distance-preserving feature extractor by adding skip connection and spectral normalization regularization, distance aware classifier using neural GP layer.
Appendix B Evaluation Metrics
In this study, we examine for each accuracy performance, general robustness performance, overconfidence performance as well as distance-awareness performance.
Robustness performance
Expected Calibration Error (abbr. ECE), Brier Score (abbr. Brier), and Negative Log Likelihood (abbr. NLL).
Overconfidence performance: OFNs% and OFPs%
Occurance of overconfident false negatives and false positives. Depending on actual applications, we may care more about Overconfident False Negatives (OFNs) in the drug cardiotoxicity task whereas in the molHIV task we care about general overconfidence, therefore both OFNs and OFPs (Overconfident False Positives). In this study, an overconfident misprediction is defined as any test sample whose predictive confidence (computed via Maximum Softmax Probability, i.e., is higher than 90% yet its prediction is wrong, namely:
where is class logit index, and is the ground truth class.
In this study we measure OFNs% and OFPs%, defined as percentages as follows:
where is -th sample’s predictive probability for positive class, i.e., .
Distance-awareness performance: DA-AUC
We design a classification task for a given test set: any molecule belonging to the close set (short distance to Train set, defined in D) gets ground-truth label 0, any molecule belonging to the far set (long distance to Train set, defined in D) gets ground truth label 1. We use predictive uncertainty (computed via 1-max) to classify if a test sample is in the close set or the far set. DA-AUC is the AUROC measurement for this task.
Appendix C Distance analysis for OFNs in CardioTox
Here we present our analysis on OFNs incurred in GNN baseline model. One of the outstanding observations is that a good portion of OFNs molecules are distant from the train set. Figure 2(a) shows a few such example molecules. More quantitatively, Figure 2(b) suggests over 60% of OFNs have Tanimoto distance , which is often regarded as a condition of having novel molecule structure. As we introduce models with more distance-awareness (e.g., GNN-GP, GNN-SNGP), distant molecules becomes less dominant in OFNs.
Appendix D Drug cardiotoxicity data split
We re-organized the original hERG (cardiotoxicity anti-target) dataset [Siramshetty et al., 2020] into CardioTox to facilitate graph model reliability research. The original dataset comes with a Train set, a compilation from ChEMBL [Gaulton et al., 2012] and NCATS (National Center for Advancing Translational Sciences), a prospective validation set from NCATS Siramshetty et al. and a test set from FDA (see Figure S3).
As a first step, we added graph structural information such as node features, edge features and adjacency matrix for each molecule record and created three test sets: Test-IID (randomly sampled 20% from the Train distribution), Test-OOD1 (the prospective validation set from NCATS) and Test-OOD2 (from FDA). The shift is mainly from input distribution (i.e., graph structures): majority of Test-IID are similar molecules to Train set, while 84% of Test-OOD1 and 82% of Test-OOD2 have novel structures.
As a second step, we further split each test set into close set and far set. Using Tanimoto graph distance [Bajusz et al., 2015] defined by molecule fingerprint representation [Rogers and Hahn, 2010], we group test samples with distance (average distance to top 8 nearest training samples) < 0.7 into close set and remaining samples go to far set. Figure S3 shows the detailed splitting flow and final counts.
Appendix E Related work
On GNN robustness research, Geisler et al. proposed a novel GNN aggregation function, Soft Medoid, to improve robustness against adversarial attacks from edge perturbation. Feng et al. constructed a two-branched GNN architecture where the first branch computes model uncertainty and data uncertainty, the second utilizes those uncertainties to adjust attention during aggregation. Those studies are organized around defending adversarial attacks, while our work aims to mitigate overconfidence issue by introducing targeted techniques.
In the area of molecule graph learning, Hwang et al. have recently applied Deep Ensemble, Monte Carlo dropout, Stochastic Gradient Langevin Dynamics (SGLD), Stochastic Weight Averaging (SWA), and Stochastic Weight Averaging Gaussian (SWAG) to GNN models and those techniques can be categorized as different approaches to create ensemble models. For example, SWA picks ensemble members along the training procedure. Our work focuses on a single GNN model and introduces architectural changes to improve its distance-awareness to reduce overconfident mispredictions. Another related work is from Hirschfeld et al. where a GNN model is first trained and its latent representation gets fed to a downstream Gaussian Processes. In contrast to that, we develop an end-to-end trainable GNN-SNGP without computational constraint of kernel computation and storage that come with exact Gaussian processes, making GNN-SNGP better suited to processing large scale datasets such as high throughput screening data in drug discovery [McCloskey et al., 2020].
On the distance-preservation/awareness side, compared with SNGP [Liu et al., 2020] techniques used in this work, two-sided gradient penalty [Gulrajani et al., 2017, Van Amersfoort et al., 2020] constrains gradient deviation by placing penalty on the loss and can be sensitive to tune. Deep Invertible Networks [Jacobsen et al., 2018] achieves high level of distance-preservation via invertiable layers but are often difficult to train and expensive in memory.
Appendix F Uncertainty improvements for OFNs
We have carefully looked into the OFNs generated by the GNN baseline, and assess their uncertainty improvements by monitoring uncertainty increase ratio (UIR):
where are the uncertainty estimates by the GNN-GP model and GNN baseline model. UIR > 1 indicates that the GNN-GP becomes less overconfident than the GNN baseline. Figure S4 lists UIRs for the 40 top overconfident false negatives from the GNN baseline. We observe over 80% of them have improved uncertainty estimates (UIR>1). As annotated in Figure S4, many highly improved OFNs are the ones distant from the Train set, while the ones getting worse uncertainty estimates often are close to the Train set. This is expected in a model with high distance-awareness.
Appendix G Results on existing benchmarks
In addition to CardioTox, we have applied our GNN-GP and GNN-SNGP models to three established molecule benchmarks: molHIV, BBBP and BACE. Specifically, molHIV is a dataset of HIV antiviral activity (each molecule has an active or inactive label), BBBP a dataset of Brain-Blood Barrier Penetration (each molecule has a label indicating whether it can penetrate through brain cell membrane to enter central nervous system) and BACE a dataset of binding affinity against human beta-secretase 1 (each molecule has a label indicating whether it binds to human beta-secretase 1). Table S1 shows the results.
Appendix H Ablation study on performance contribution
Distance-preservation in representation learning and distance-awareness in classifier both impact the performances we care about in this study. In order to understand relative contributions of them to the performance improvement, we further carried out an extensive ablation study. We take the latent representations learned by GNN baseline, GNN-GP and GNN-SNGP (increasing distance-preservation) and feed them to classifiers with increasing levels of distance-awareness: Dense layer (using logistic regression), GP-layer, exact Gaussian processes classifier (GPC hereafter). We also added experiments using pre-defined representation (i.e., molecule fingerprint denoted as FPs) which should provide 100% of distance preservation. Table S2 shows the results.
We find a GNN model achieves best performances when both distance-preservation (via SN) and distance-awareness (via GP) are present in the architecture (i.e., GNN-SNGP embeddings with GPC). However, one interesting note is that naively increasing distance-preservation does not guarantee generalization, either in-domain or OOD. For example, we observe that the pre-defined representation (i.e., molecule fingerprint denoted as FPs) offers perfect distance-preservation, but tends to give relatively low AUROC (compared with neural representation) under data shifts such as Test-OOD2 even with exact GPC as its classifier. We hypothesis that this is related to a general trade-off in representation learning between dimension reduction and information preservation. The dimension reduction aspect tries to discard information noisy and/or less relevant to prediction task so that given limited data, models could escape the curse of dimensionality and achieve good performance under finite data. However, this could also lead to the feature collapse phenomenon that harms model robustness especially under distributional shifts. On the other hand, the information preservation aspect recovers robustness by mitigating feature collapse, but keeping around noisy/irrelevant features may demand more data for convergence and thus negatively impact accuracy given limited data. This work tries to find a good balance between these two aspects by designing a distance-preserving representation that can be optimized toward the task at hand.