Pre-training of Graph Augmented Transformers for Medication Recommendation

Junyuan Shang, Tengfei Ma, Cao Xiao, Jimeng Sun

Introduction

The availability of massive electronic health records (EHR) data and the advances of deep learning technologies have provided unprecedented resource and opportunity for predictive healthcare, including the computational medication recommendation task. A number of deep learning models were proposed to assist doctors in making medication recommendation Xiao et al. 2018a; Shang et al. 2019; Baytas et al. 2017; Choi et al. 2018; Ma et al. 2018. They often learn representations for medical entities (e.g., patients, diagnosis, medications) from patient EHR data, and then use the learned representations to predict medications that are suited to the patient’s health condition.

To provide effective medication recommendation, it is important to learn accurate representation of medical codes. Despite that various considerations were handled in previous works for improving medical code representations Ma et al. 2018; Baytas et al. 2017; Choi et al. 2018, there are two limitations with the existing work:

Selection bias: Data that do not meet training data criteria are discarded before model training. For example, a large number of patients who only have one hospital visit were discarded from training in Shang et al. 2019.

Lack of hierarchical knowledge: For medical knowledge such as diagnosis code ontology (Figure. 1), their internal hierarchical structures were rarely embedded in their original graph form when incorporated into representation learning.

To mitigate the aforementioned limitations, we propose G-BERT that combines the pre-training techniques and graph neural networks for better medical code representation and medication recommendation. G-BERT is enabled and demonstrated by the following technical contributions:

Pre-training to leverage more data: Pre-training techniques, such as ELMo Peters et al. 2018, OpenAI GPT Radford et al. 2018 and BERT Devlin et al. 2018, have demonstrated a notably good performance in various natural language processing tasks. These techniques generally train language models from unlabeled data, and then adapt the derived representations to different tasks by either feature-based (e.g. ELMo) or fine-tuning (e.g. OpenAI GPT, BERT) methods. We developed a new pre-training method based on BERT for pre-training on each visit of EHR so that the data with only one hospital visit can also be utilized. We revised BERT to fit EHR data in both input and pre-training objectives. To our best knowledge, G-BERT is the first model that leverages Transformers and language model pre-training techniques in healthcare domain. Compared with other supervised models, G-BERT can utilize discarded/unlabeled data more efficiently.

Medical ontology embedding with graph neural networks: We enhance the representation of medical codes via learning medical ontology embedding for each medical codes with graph neural networks. We then input the ontology embedding into a multi-layer Transformer Vaswani et al. 2017 for BERT-style pre-training and fine-tuning.

Related Work

Medication Recommendation can be categorized into instance-based and longitudinal recommendation methods Shang et al. 2019. Instance-based methods focus on current health conditions. Among them, Leap Zhang et al. 2017 formulates a multi-instance multi-label learning framework and proposes a variant of sequence-to-sequence model based on content-attention mechanism to predict combination of medicines given patient’s diagnoses. Longitudinal-based methods leverage the temporal dependencies among clinical events, see Choi et al. 2016; Xiao et al. 2018b; Lipton et al. 2015. Among them, RETAIN Choi et al. 2016 uses a two-level neural attention model to detect influential past visits and significant clinical variables within those visits for improved medication recommendation.

The goal of pre-training techniques is to provide model training with good initializations. Pre-training has been shown extremely effective in various areas such as image classification Hinton et al. 2006 and machine translation Ramachandran et al. 2016. The unsupervised pre-training can be considered as a regularizer that supports better generalization from the training dataset Erhan et al. 2010. Recently, language model pre-training techniques such as Peters et al. 2018; Radford et al. 2018; Devlin et al. 2018 have shown to largely improve the performance on multiple NLP tasks. As the most widely used one, BERT Devlin et al. 2018 builds on the Transformer Vaswani et al. 2017 architecture and improves the pre-training using a masked language model for bidirectional representation. In this paper, we adapt the framework of BERT and pre-train our model on each visit of the EHR data to leverage the single-visit data that were not fit for training in other medication recommendation models.

GNNs are neural networks that learn node or graph representations from graph-structured data. Various graph neural networks have been proposed to encode the graph-structure information, including graph convolutional neural networks (GCN) Kipf and Welling 2017, message passing networks (MPNN) Gilmer et al. 2017, graph attention networks (GAT) Velickovic et al. 2017. GNNs have already been demonstrated useful on EHR modeling Choi et al. 2017; Shang et al. 2019. GRAM Choi et al. 2017 represented a medical concept as a combination of its ancestors in the medical ontology using an attention mechanism. It’s different from G-BERT from two aspects as described in Section 4.2. Another work worth mentioning is GAMENet Shang et al. 2019, which also used graph neural network to assist the medication recommendation task. However, GAMENet has a different motivation which results in using graph neural networks on drug-drug-interaction graphs instead of medical ontology.

Problem Formalization

Medical codes are usually categorized according to a tree-structured classification system such as ICD-9 ontoloy for diagnosis and ATC ontology for medication. We use Od,Om\mathcal{O}_{d},\mathcal{O}_{m} to denote the ontology for diagnosis and medication. Similarly, we use O∗\mathcal{O}_{\ast} to indicate the unified definition for different type of medical codes. In detial, O∗=C∗‾∪C∗\mathcal{O}_{\ast}=\overline{\mathcal{C}_{\ast}}\cup\mathcal{C}_{\ast} where C∗‾\overline{\mathcal{C}_{\ast}} denotes the codes excluding leaf codes. For simplicity, we define two function pa(⋅),ch(⋅)pa(\cdot),ch(\cdot) which accept target medical code and return ancestors’ code set and direct child code set.

Given diagnosis codes Cdt\mathcal{C}_{d}^{t} of the visit at time tt, patient history X1:t={X1,X2,⋯ ,Xt−1}\mathcal{X}_{1:t}=\{\mathcal{X}_{1},\mathcal{X}_{2},\cdots,\mathcal{X}_{t-1}\}, we want to recommend multiple medications by generating multi-label output y^t∈{0,1}∣Cm∣\hat{\bm{y}}_{t}\in\{0,1\}^{|\mathcal{C}_{m}|}.

Method

The overall framework of G-BERT is described in Figure 2. G-BERT first derives the initial embedding of medical codes from medical ontology using graph neural networks. Then, in order to fully utilize the rich EHR data, G-BERT constructs an adaptive BERT model on the discarded single-visit data for visit representation. Finally we add a prediction layer and fine-tune the model in the medication recommendation task. In the following we will describe G-BERT in detail. But firstly, we give a brief background of BERT especially for the two pre-training objectives which will be later adapted to EHR data in Section 4.3.

Based on a multi-layer Transformer encoder Vaswani et al. 2017 (The transformer architecture has been ubiquitously used in many sequence modeling tasks recently, so we will not introduce the details here), BERT is pre-trained using two unsupervised tasks:

Masked Language Model. Instead of predicting words based on previous words, BERT randomly selects words to mask out and then predicts the original vocabulary IDs of the masked words from their (bidirectional) context.

Next Sentence Prediction. Many of BERT’s downstream tasks are predicting the relationships of two sentences, thus in the pre-training phase, BERT has am a binary sentence prediction task to predict whether one sentence is the next sentence of the other.

A typical input to BERT is as follows ( Devlin et al. 2018):

Input = [CLS] the man went to [MASK] store [SEP] he bought a gallon [MASK] milk [SEP] Label = IsNext

where [CLS] is the first token of each sentence pair to represent the special classification embedding, i.e. the final state of this token is used as the aggregated sequence representation for classification tasks; [SEP] is used to separate two sentences; [MASK] is used to mask out the predicted words in the masked language model. Using this form, these inputs facilitate the two tasks described above, and they will also be used in our method description in the following section.

2 Input Representation

The G-BERT model takes medical codes’ ontology embeddings as input, and obtains intermediate representations from a Transformer encoder as the visit embeddings. It is then pre-trained on EHR from patients who only have one hospital visit. The derived encoder and visit embedding will be fed into a classifier and fine-tuned to make predictions.

We constructed ontology embedding from diagnosis ontology Od\mathcal{O}_{d} and medication ontology Om\mathcal{O}_{m}. Since the medical codes in raw EHR data can be considered as leaf nodes in these ontology trees, we can enhance the medical code embedding using graph neural networks (GNNs) to integrate the ancestors’ information of these codes. Here we perform a two-stage procedure with a specially designed GNN for ontology embedding.

where g(⋅,⋅,⋅)g(\cdot,\cdot,\cdot) is an aggregation function which accepts the target medical code c∗c_{\ast}, its direct child codes ch(c∗)ch(c_{\ast}) and initial embedding matrix. Intuitively, the aggregation function can pass and fuse information in target node from its direct children which result in the more related embedding of ancestor’ code to child codes’ embedding.

where g(⋅,⋅,⋅)g(\cdot,\cdot,\cdot) accepts ancestor codes of target medical code c∗c_{\ast}. Here, we use pa(c∗)pa(c_{\ast}) instead of ch(c∗)ch(c_{\ast}), since utilizing the ancestors’ embedding can indirectly associate all medical codes instead of taking each leaf code as independent input.

The option for the aggregation function g(⋅,⋅,⋅)g(\cdot,\cdot,\cdot) is flexible, including sum, mean. Here we choose the one from graph attention networks (GAT) Velickovic et al. 2017, which has shown efficient embedding learning ability on graph-structured tasks, e.g., node classification and link prediction. In particular, we implement the aggregation function g(⋅,⋅,⋅)g(\cdot,\cdot,\cdot) as follows (taking stage 2 for an example):

As shown in Figure 2, we construct ICD-9 tree for diagnosis and ATC tree for medication using the same structure. Here the direction of arrow shows the information flow where ancestor nodes can get information from their direct children (in stage 1) and similarly leaf nodes can get information from their connected ancestors (in stage 2).

It is worth mentioning that our graph embedding method on medical ontology is different from GRAM Choi et al. 2017 from the following two aspects:

Initialization: we initialize all the node embeddings from a learnable embedding matrix, while GRAM learns them using Glove from the co-occurrence information.

Updating: we develop a two-step updating function for both leaf nodes and ancestor nodes; while in GRAM, only the leaf nodes are updated (as a combination of their ancestor nodes and themselves).

2.2 Visit Embedding

where [CLS] is a special token as in BERT. It is put in the first position of each visit of type ∗\ast and its final state can be used as the representation of the visit. Intuitively, it is more reasonable to use Transformers as encoders (multi-head attention based architecture) than RNN or mean/sum to aggregate multiple medical embedding for visit embedding since the set of medical codes within one visit is not ordered. Note that symbol [SEP] is also ignored considering there is no clear separate among codes within one visit.

It is worth noting that our Transformer encoder is different from the original one in the position embedding part. Position embedding, as an important component in Transformers and BERT, is used to encode the position and order information of each token in a sequence. However, one big difference between language sentences and EHR sequences is that the medical codes within the same visit do not generally have an order, so we remove the position embedding in our model.

3 Pre-training

Thus, for the self-prediction task, we want the visit embedding v∗v_{\ast} to recover what it is made of, i.e., the input medical codes C∗\mathcal{C}_{\ast} limited by the same type for each visit as follows:

where C∗(n)\mathcal{C}_{\ast}^{(n)} is the medical codes set of nn-th patient, ∗∈{d,m}*\in\{d,m\} and we minimize the binary cross entropy loss Lse\mathcal{L}_{se}. For instance, assume that the nn-th patient takes 10 different medications out of total 100 medications which means ∣Cm(n)∣=10|\mathcal{C}_{m}^{(n)}|=10 and ∣Cm∣=100|\mathcal{C}_{m}|=100. In such case, we instantiate and minimize Lse(vm,Cm(n))L_{se}(\bm{v}_{m},\mathcal{C}_{m}^{(n)}) to produce high probabilities among 10 taken medications captured by −∑c∗∈Cm(n)log⁡p(c∗∣vm)-\sum_{c_{\ast}\in\mathcal{C}_{m}^{(n)}}\log p(c_{\ast}|\bm{v}_{m}) and lower the probabilities among 90 non-taken ones captured by ∑c∗∈{Cm∖Cm(n)}log⁡p(c∗∣vm)\sum_{c_{\ast}\in\{\mathcal{C}_{m}\setminus\mathcal{C}_{m}^{(n)}\}}\log p(c_{\ast}|\bm{v}_{m}). In practise, Sigmoid(f(v∗))\text{Sigmoid}(f(\bm{v}_{\ast})) should be transformed by applying a fully connected neural network f(⋅)f(\cdot) with one hidden layer. With an analogy to the Masked LM task in BERT, we also used specific symbol [MASK] to randomly replace the original medical code c∗∈C∗c_{\ast}\in\mathcal{C}_{\ast}. So there are 15%15\% codes in C∗\mathcal{C}_{\ast} which will be replaced randomly and the model should have the ability to predict the masked code based on others.

Likewise, for the dual-prediction task, since the visit embedding v∗\bm{v}_{\ast} carries the information of medical codes of type ∗\ast, we can further expect it has the ability to do more task-specific prediction as follows:

where we use the same transformation function Sigmoid(f1(vm))\text{Sigmoid}(f_{1}(\bm{v}_{m})), Sigmoid(f2(vd))\text{Sigmoid}(f_{2}(\bm{v}_{d})) f1,f2f_{1},f_{2} are the multiple layer perceptron (MLP) with one hidden layer with different weight matrix to transform the visit embedding and optimize the binary cross entropy loss Ldu\mathcal{L}_{du} expanded same as Lse\mathcal{L}_{se} in Eq. 6. This is a direct adaptation of the next sentence prediction task. In BERT, the next sentence prediction task facilitates the prediction of sentence relations, which is a common task in NLP. However, in healthcare, most predictive tasks do not have a sequence pair to classify. Instead, we are often interested in predicting unknown disease or medication codes of the sequence. For example, in medication recommendation, we want to predict multiple medications given only the diagnosis codes. Inversely, we can also predict unknown diagnosis given the medication codes.

Thus, our final pre-training optimization objective can simply be the combination of the aforementioned losses, as shown in Eq. 8.

In practise, we could integrate using mini-batch technique to train on EHR data from all patients with a single visit.

4 Fine-tuning

After obtaining pre-trained visit representation for each visit, for a prediction task on a multi-visit sequence data, we aggregate all the visit embedding and add a prediction layer for the medication recommendation task as shown in Figure. 4. To be specific, from pre-training on all visits, we have a pre-trained Transformer encoder, which can then be used to get the visit embedding v∗τ\bm{v}_{*}^{\tau} at time τ\tau. The known diagnosis codes Cdt\mathcal{C}_{d}^{t} at the prediction time tt is also represented using the same model as v∗t\bm{v}_{*}^{t}. Concatenating the mean of previous diagnoses visit embeddings and medication visit embeddings, also the last diagnoses visit embedding, we built an MLP based prediction layer to predict the recommended medication codes as in Equation 9.

Given the true labels y^t\hat{\bm{y}}_{t} at each time stamp tt, the loss function for the whole EHR sequence (i.e. a patient) is

Experiment

We used EHR data from MIMIC-III Johnson et al. 2016 and conducted all our experiments on a cohort where patients have more than one visit. We utilize data from patients with both single visit and multiple visits in the training dataset as pre-training data source (multi-visit data are split into visit slices and duplicate codes within a single visit are removed in order to avoid leakage of information). In this work, we transform the drug coding from NDC to ATC Third Level for using the ontology information. The statistics of the datasets are summarized in Table 2.

1.2 Baselines

We compared G-BERT https://github.com/jshang123/G-Bert with the following baselines. All methods are implemented in PyTorch Paszke et al. 2017 and trained on an Ubuntu 16.04 with 8GB memory and Nvidia 1080 GPU.

Logistic Regression (LR) is logistic regression with L1/L2 regularization. Here we represent sequential multiple medical codes by sum of multi-hot vector of each visit. Binary relevance technique Luaces et al. 2012 is used to handle multi-label output.

LEAP Zhang et al. 2017 is an instance-based medication combination recommendation method which formalizes the task in multi-instance and multi-label learning framework. It utilizes a encoder-decoder based model with attention mechanism to build complex dependency among diseases and medications.

RETAIN Choi et al. 2016 makes sequential prediction of medication combination and diseases prediction based on a two-level neural attention model that detects influential past visits and clinical variables within those visits.

GRAM Choi et al. 2017 injects domain knowledge (ICD9 Dx code tree) to tanh via attention mechanism.

GAMENet Shang et al. 2019 is the method to recommend accuracy and safe medication based on memory neural networks and graph convolutional networks by leveraging EHR data and Drug-Drug Interaction (DDI) data source. For fair comparison, we use a variant of GAMENet without DDI knowledge and procedure codes as input renamed as GAMENet−\text{GAMENet}^{-}.

G-BERT is our proposed model which integrated the GNN representation into Transformer-based visit encoder with pre-training on single-visit EHR data.

We also evaluated 3 G-BERT variants for model ablation.

G-BERTG−,P−\texttt{G-BERT}_{G^{-},P^{-}}: We directly use medical embedding without ontology information as input and initialize the model’s parameters without pre-training.

G-BERTG−\texttt{G-BERT}_{G^{-}}: We directly use medical embedding without ontology information as input with pre-training.

G-BERTP−\texttt{G-BERT}_{P^{-}}: We use ontology information to get ontology embedding as input and initialize the model’s parameters without pre-training.

1.3 Metrics

To measure the prediction accuracy, we used Jaccard Similarity Score (Jaccard), Average F1 (F1) and Precision Recall AUC (PR-AUC). Jaccard is defined as the size of the intersection divided by the size of the union of ground truth set Yt(k)Y_{t}^{(k)} and predicted set Y^t(k)\hat{Y}_{t}^{(k)}.

where NN is the number of patients in test set and TkT_{k} is the number of visits of the kthk^{th} patient.

1.4 Implementation Details

We randomly divide the dataset into training, validation and testing set in a 0.6:0.2:0.20.6:0.2:0.2 ratio. For G-BERT, the hyperparameters are adjusted on evaluation set: (1) GAT part: input embedding dimension as 75, number of attention heads as 4; (2) BERT part: hidden dimension as 300, dimension of position-wise feed-forward networks as 300, 2 hidden layers with 4 attention heads for each layer. Specially, we alternated the pre-training with 5 epochs and fine-tuning procedure with 5 epochs for 15 times to stabilize the training procedure.

For LR, we use the grid search over typical range of hyper-parameter to search the best hyperparameter values which result in L1 norm penalty with weight as 1.11.1. For deep learning models, we implemented RNN using a gated recurrent unit (GRU) Cho et al. 2014 and utilize dropout with a probability of 0.4 on the output of embedding. We test several embedding choice for baseline methods and determine the dimension for medical embedding as 300 and thershold for final prediction as 0.3 for better performance. Training is done through Adam Kingma and Ba 2014 at learning rate 5e-4. We fix the best model on evaluation set within 100 epochs and report the performance in test set.

2 Results

Table. 3 compares the performance on the medication recommendation task. For variants of G-BERT, G-BERTG−,P−\texttt{G-BERT}_{G^{-},P^{-}} performs worse compared with G-BERTG−\texttt{G-BERT}_{G^{-}} and G-BERTP−\texttt{G-BERT}_{P^{-}} which demonstrate the effectiveness of using ontology information to get enhanced medical embedding as input and employ an unsupervised pre-training procedure on larger abundant data. Incorporating both hierarchical ontology information and pre-training procedure, the end-to-end model G-BERT has more capacity and achieve comparable results with others.

As for baseline models, LR and Leap are worse than our most basic model (G-BERTG−,P−\texttt{G-BERT}_{G^{-},P^{-}}) in terms of most metrics. Comparing G-BERTP−\texttt{G-BERT}_{P^{-}} and GRAM, which both used medical ontology information without pre-training, the scores of our G-BERTP−\texttt{G-BERT}_{P^{-}} is slightly higher in all metrics. This can demonstrate the validness of using Transformer encoders and the specific prediction layer for medication recommendation. Our final model G-BERT is also better than the attention based model, RETAIN, and the recently published state-of-the-art model, GAMENet. Specifically, even adding the extra information of DDI knowledge and procedure codes, GAMENet still performs worse than G-BERT.

In addition, we visualized the pre-training medical code embeddings of G-BERTG−\texttt{G-BERT}_{G^{-}} and G-BERT to show the effectiveness of ontology embedding using online embedding projector https://projector.tensorflow.org/ shown in (https://raw.githubusercontent.com/jshang123/G-Bert/master/saved/tsne.png/).

Conclusion

In this paper we proposed a pre-training model named G-BERT for medical code representation and medication recommendation. To our best knowledge, G-BERT is the first that utilizes language model pre-training techniques in healthcare domain. It adapted BERT to the EHR data and integrated medical ontology information using graph neural networks. By additional pre-training on the EHR from patients who only have one hospital visit which are generally discarded before model training, G-BERT outperforms all baselines in prediction accuracy on medication recommendation task. One direction for the future work is to add more auxiliary and structural tasks to improve the ability of code representaion. Another direction may be to adapt our model to be suitable for even larger datasets with more heterogeneous modalities.

Acknowledgments

This work was supported by the National Science Foundation award IIS-1418511, CCF-1533768 and IIS-1838042, the National Institute of Health award 1R01MD011682-01 and R56HL138415.

References