GRAM: Graph-based Attention Model for Healthcare Representation Learning

Edward Choi, Mohammad Taha Bahadori, Le Song, Walter F. Stewart, Jimeng Sun

Introduction

The rapid growth in volume and diversity of health care data from electronic health records (EHR) and other sources is motivating the use of predictive modeling to improve care for individual patients. In particular, novel applications are emerging that use deep learning methods such as word embedding (Choi et al., 2016c; Choi et al., 2016d), recurrent neural networks (RNN) (Che et al., 2016; Choi et al., 2016a, b; Lipton et al., 2016), convolutional neural networks (CNN) (Nguyen et al., 2016) or stacked denoising autoencoders (SDA) (Che et al., 2015; Miotto et al., 2016), demonstrating significant performance enhancement for diverse prediction tasks. Deep learning models appear to perform significantly better than logistic regression or multilayer perceptron (MLP) models that depend, to some degree, on expert feature construction (Lipton et al., 2015; Razavian et al., 2016).

Training deep learning models typically requires large amounts of data that often cannot be met by a single health system or provider organization. Sub-optimal model performance can be particularly challenging when the focus of interest is predicting onset of a rare disease. For example, using Doctor AI (Choi et al., 2016a), we discovered that RNN alone was ineffective to predict the onset of diseases such as cerebral degenerations (e.g. Leukodystrophy, Cerebral lipidoses) or developmental disorders (e.g. autistic disorder, Heller’s syndrome), partly because their rare occurrence in the training data provided little learning opportunity to the flexible models like RNN.

The data requirement of deep learning models comes from having to assess exponential number of combinations of input features. This can be alleviated by exploiting medical ontologies that encodes hierarchical clinical constructs and relationships among medical concepts. Fortunately, there are many well-organized ontologies in healthcare such as the International Classification of Diseases (ICD), Clinical Classifications Software (CCS) (Stearns et al., 2001) or Systematized Nomenclature of Medicine-Clinical Terms (SNOMED-CT) (Project et al., 2010). Nodes (i.e. medical concepts) close to one another in medical ontologies are likely to be associated with similar patients, allowing us to transfer knowledge among them. Therefore, proper use of medical ontologies will be helpful when we lack enough data for the nodes in the ontology to train deep learning models.

In this work, we propose GRAM, a method that infuses information from medical ontologies into deep learning models via neural attention. Considering the frequency of a medical concept in the EHR data and its ancestors in the ontology, GRAM decides the representation of the medical concept by adaptively combining its ancestors via attention mechanism. This will not only support deep learning models to learn robust representations without large amount of data, but also learn interpretable representations that align well with the knowledge from the ontology. The attention mechanism is trained in an end-to-end fashion with the neural network model that predicts the onset of disease(s). We also propose an effective initialization technique in addition to the ontological knowledge to better guide the representation learning process.

We compare predictive performance (i.e. accuracy, data needs, interpretability) of GRAM to various models including the recurrent neural network (RNN) in two sequential diagnoses prediction tasks and one heart failure (HF) prediction task. We demonstrate that GRAM is up to 10% more accurate than the basic RNN for predicting diseases less observed in the training data. After discussing GRAM’s scalability, we visualize the representations learned from various models, where GRAM provides more intuitive representations by grouping similar medical concepts close to one another. Finally, we show GRAM’s attention mechanism can be interpreted to understand how it assigns the right amount of attention to the ancestors of each medical concept by considering the data availability and the ontology structure.

Methodology

We first define the notations describing EHR data and medical ontologies, followed by a description of GRAM (Section 2.2), the end-to-end training of the attention generation and predictive modeling (Section 2.3), and the efficient initialization scheme (Section 2.4).

We denote the set of entire medical codes from the EHR as c1,c2,…,c_{1},c_{2},\ldots, c∣C∣∈Cc_{|\mathcal{C}|}\in\mathcal{C} with the vocabulary size ∣C∣|\mathcal{C}|. The clinical record of each patient can be viewed as a sequence of visits V1,…,VTV_{1},\ldots,V_{T} where each visit contains a subset of medical codes Vt⊆CV_{t}\subseteq\mathcal{C}. VtV_{t} can be represented as a binary vector xt∈{0,1}∣C∣\mathbf{x}_{t}\in\{0,1\}^{|\mathcal{C}|} where the ii-th element is 1 only if VtV_{t} contains the code cic_{i}. To avoid clutter, all algorithms will be presented for a single patient.

We assume that a given medical ontology G\mathcal{G} typically expresses the hierarchy of various medical concepts in the form of a parent-child relationship, where the medical codes C\mathcal{C} form the leaf nodes. Ontology G\mathcal{G} is represented as a directed acyclic graph (DAG) whose nodes form a set D=C+C′\mathcal{D}=\mathcal{C}+\mathcal{C^{\prime}}. The set C′={c∣C∣+1,c∣C∣+2,…,c∣C∣+∣C′∣}\mathcal{C^{\prime}}=\{c_{|\mathcal{C}|+1},c_{|\mathcal{C}|+2},\ldots,c_{|\mathcal{C}|+|\mathcal{C^{\prime}}|}\} consists of all non-leaf nodes (i.e. ancestors of the leaf nodes), where ∣C′∣|\mathcal{C^{\prime}}| represents the number of all non-leaf nodes. We use knowledge DAG to refer to G\mathcal{G}. A parent in the knowledge DAG G\mathcal{G} represents a related but more general concept over its children. Therefore, G\mathcal{G} provides a multi-resolution view of medical concepts with different degrees of specificity. While some ontologies are exclusively expressed as parent-child hierarchies (e.g. ICD-9, CCS), others are not. For example, in some instances SNOMED-CT also links medical concepts to causal or treatment relationships, but the majority relationships in SNOMED-CT are still parent-child. Therefore, we focus on the parent-child relationships in this work.

2 Knowledge DAG and the Attention Mechanism

GRAM leverages the parent-child relationship of G\mathcal{G} to learn robust representations when data volume is constrained. GRAM balances the use of ontology information in relation to data volume in determining the level of specificity for a medical concept. When a medical concept is less observed in the data, more weight is given to its ancestors as they can be learned more accurately and offer general (coarse-grained) information about their children. The process of resorting to the parent concepts can be automated via the attention mechanism and the end-to-end training as described in Figure 1.

f(ei,ej)f(\mathbf{e}_{i},\mathbf{e}_{j}) is a scalar value representing the compatibility between the basic embeddings of ei\mathbf{e}_{i} and ek\mathbf{e}_{k}. We compute f(ei,ej)f(\mathbf{e}_{i},\mathbf{e}_{j}) via the following feed-forward network with a single hidden layer (MLP),

Remarks: The example in Figure 1 is derived based on a single path from cic_{i} to cac_{a}. However, the same mechanism can be applicable to multiple paths as well. For example, code ckc_{k} has two paths to the root cac_{a}, containing five ancestors in total. Another scenario is where the EHR data contain both leaf codes and some ancestor codes. We can move those ancestors present in EHR data from the set C′\mathcal{C^{\prime}} to C\mathcal{C} and apply the same process as Eq. (1) to obtain the final representations for them.

3 End-to-End Training with a Predictive Model

We train the attention mechanism together with a predictive model such that the attention mechanism improves the predictive performance. By concatenating final representation g1,g2,…,g∣C∣\mathbf{g}_{1},\mathbf{g}_{2},\ldots,\mathbf{g}_{|\mathcal{C}|} of all medical codes, we have the embedding matrix G∈Rm×∣C∣\mathbf{G}\in\mathcal{R}^{m\times|\mathcal{C}|} where gi\mathbf{g}_{i} is its ii-th column of G\mathbf{G}. We can then convert visit VtV_{t} to a visit representation vt\mathbf{v}_{t} by multiplying embedding matrix G\mathbf{G} with multi-hot vector xt\mathbf{x}_{t} indicating the clinical events in visit VtV_{t} as shown in the right side of Figure 1. Finally the visit representation vt\mathbf{v}_{t} will be used as an input to pass to a predictive model for predicting the target label yt\mathbf{y}_{t} using a neural network (NN) model. In this work, we use RNN as the choice of the NN model as the task is to perform sequential diagnoses prediction (Choi et al., 2016a, b) with the objective of predicting the disease codes of the next visit Vt+1V_{t+1} given the visit records up to the current timestep V1,V2,…,VtV_{1},V_{2},\ldots,V_{t}, which can be expressed as follows,

where we sum the cross entropy errors from all timestamps of y^t\widehat{\mathbf{y}}_{t}, TT denotes the number of timestamps of the visit sequence. Note that the above loss is defined for a single patient. In actual implementation, we will take the average of the individual loss for multiple patients. Algorithm 1 describes the overall training procedure of GRAM, under the assumption that we are performing the sequential diagnoses prediction task using an RNN. Note that Algorithm 1 describes stochastic gradient update to avoid clutter, but it can be easily extended to other gradient based optimization such as mini-batch gradient update.

4 Initializing Basic Embeddings

The attention generation mechanism in Section 2.2 requires basic embeddings ei\mathbf{e}_{i} of each node in the knowledge DAG. The basic embeddings of ancestors, however, pose a difficulty because they are often not observed in the data. To properly initialize them, we use co-occurrence information to learn the basic embeddings of medical codes and their ancestors. Co-occurrence has proven to be an important source of information when learning representations of words or medical concepts (Mikolov et al., 2013; Choi et al., 2016c; Choi et al., 2016d). To train the basic embeddings, we employ GloVe (Pennington et al., 2014), which uses the global co-occurrence matrix of words to learn their representations. In our case, the co-occurrence matrix of the codes and the ancestors was generated by counting the co-occurrences within each visit VtV_{t}, where we augment each visit with the ancestors of the codes in the visit.

We describe the details of the initialization algorithm with an example. We borrow the parent-child relationships from the knowledge DAG of Figure 1. Given a visit VtV_{t},

we augment it with the ancestors of all the codes to obtain the augmented visit Vt′V^{\prime}_{t},

where the augmented ancestors are underlined. Note that a single ancestor can appear multiple times in Vt′V^{\prime}_{t}. In fact, the higher the ancestor is in the knowledge DAG, the more times it is likely to appear in Vt′V^{\prime}_{t}. We count the co-occurrence of two codes in Vt′V^{\prime}_{t} as follows,

where the hyperparameters xmaxx_{max} and α\alpha are respectively set to 100100 and 0.750.75 as the original paper (Pennington et al., 2014). Note that, after the initialization, the basic embeddings ei\mathbf{e}_{i}’s of both leaf nodes (i.e. medical codes) and non-leaf nodes (i.e. ancestors) are fine-tuned during model training via backpropagation.

Experiments

We conduct three experiments to determine if GRAM offered superior prediction performance when facing data insufficiency. We first describe the experimental setup followed by results comparing predictive performance of GRAM with various baseline models. After discussing GRAM’s scalability, we qualitatively evaluate the interpretability of the resulting representation. The source code of GRAM is publicly available at https://github.com/mp2893/gram.

Prediction tasks and source of data: We conduct the sequential diagnoses prediction (SDP) tasks on two datasets, which aim at predicting all diagnosis categories in the next visit, and a heart failure (HF) prediction task on one dataset, which is a binary prediction task for predicting a future HF onset where the prediction is made only once at the last visit xT\mathbf{x}_{T}. Two sequential diagnoses predictions (SDP) are respectively conducted using two datasets: 1) Sutter Palo Alto Medical Foundation (PAMF) dataset, which consists of 18-years longitudinal medical records of 258K patients between age 50 and 90. This will determine GRAM’s performance for general adult population with long visit records. 2) MIMIC-III dataset (Johnson et al., 2016; Goldberger et al., 2000), which is a publicly available dataset consisting of medical records of 7.5K intensive care unit (ICU) patients over 11 years. This will determine GRAM’s performance for high-risk patients with very short visit records. We utilize all the patients with at least 2 visits. We prepared the true labels yt\mathbf{y}_{t} by grouping the ICD9 codes into 283 groups using CCS single-level diagnosis grouperhttps://www.hcup-us.ahrq.gov/toolssoftware/ccs/AppendixASingleDX.txt. This is to improve the training speed and predictive performance for easier analysis, while preserving sufficient granularity for each diagnosis. Each diagnosis code’s varying frequency in the training data can be viewed as different degrees of data insufficiency. We calculate Accuracy@k for each of CCS single-level diagnosis codes such that, given a visit VtV_{t}, we get 1 if the target diagnosis is in the top kk guesses and 0 otherwise. We conduct HF prediction on Sutter heart failure (HF) cohort, which is a subset of Sutter PAMF data for a heart failure onset prediction study with 3.4K HF cases chosen by a set of criteria described in Vijayakrishnan et al. (2014); Gurwitz et al. (2013) and 27K matching controls chosen by a set of criteria described in Choi et al. (2016e). This will determine GRAM’s performance for a different prediction task where we predict the onset of one specific condition. We randomly downsample the training data to create different degrees of data insufficiency. We use area under the ROC curve (AUC) to measure the performance. A summary of the datasets are provided in Table 1.We used CCS multi-level diagnoses hierarchyhttps://www.hcup-us.ahrq.gov/toolssoftware/ccs/AppendixCMultiDX.txt as our knowledge DAG G\mathcal{G}. We also tested the ICD9 code hierarchyhttp://www.icd9data.com/2015/Volume1/default.htm, but the performance was similar to using CCS multi-level hierarchy. For all three tasks, we randomly divide the dataset into the training, validation and test set by .75:.10:.15 ratio, and use the validation set to tune the hyper-parameters. Further details regarding the hyper-parameter tuning are provided below. The test set performance is reported in the paper.

Implementation details: We implemented GRAM with Theano 0.8.2 (Team, 2016). For training models, we used Adadelta (Zeiler, 2012) with a mini-batch of 100 patients, on a machine equipped with Intel Xeon E5-2640, 256GB RAM, four Nvidia Titan X’s and CUDA 7.5.

Models for comparison are the following. The first two GRAM+ and GRAM are the proposed methods and the rest are baselines. Hyper-parameter tuning is configured so that the number of parameters for the baselines would be comparable to GRAM’s. Further details are provided below.

GRAM: Input sequence x1,…,xT\mathbf{x}_{1},\ldots,\mathbf{x}_{T} is first transformed by the embedding matrix G\mathbf{G}, then fed to the GRU with a single hidden layer, which in turn makes the prediction, as described by Eq. (4). The basic embeddings ei\mathbf{e}_{i}’s are randomly initialized.

GRAM+: We use the same setup as GRAM, but the basic embeddings ei\mathbf{e}_{i}’s are initialized according to Section 2.4.

RandomDAG: We use the same setup as GRAM, but each leaf concept has five randomly assigned ancestors from the CCS multi-level hierarchy to test the effect of correct domain knowledge.

RNN+: We use the RNN model with the same setup as before, but we initialize the embedding matrix Wemb\mathbf{W}_{emb} with GloVe vectors trained only with the co-occurrence of leaf concepts. This is to compare GRAM with a similar weight initialization technique.

SimpleRollUp: We use the RNN model with the same setup as before. But for input xt\mathbf{x}_{t}, we replace all diagnosis codes with their direct parent codes in the CCS multi-level hierarchy, giving us 578, 526 and 517 input codes respectively for Sutter data, MIMIC-III and Sutter HF cohort. This is to compare the performance of GRAM with a common grouping technique.

RollUpRare: We use the RNN model with the same setup as before, but we replace any diagnosis code whose frequency is less than a certain threshold in the dataset with its direct parent. We set the threshold to 100 for Sutter data and Sutter HF cohort, and 10 for MIMIC-III, giving us 4,408, 935 and 1,538 input codes respectively for Sutter data, MIMIC-III and Sutter HF cohort. This is an intuitive way of dealing with infrequent medical codes.

Hyper-parameter Tuning: We define five hyper-parameters for GRAM:

dimensionality mm of the basic embedding ei\mathbf{e}_{i}:

dimensionality rr of the RNN hidden layer ht\mathbf{h}_{t} from Eq. (4):

dimensionality ll of Wa\mathbf{W}_{a} and ba\mathbf{b}_{a} from Eq. (3):

L2L_{2} regularization coefficient for all weights except RNN weights: [0.1, 0.01, 0.001, 0.0001]

dropout rate for the dropout on the RNN hidden layer: [0.0, 0.2, 0.4, 0.6, 0.8]

We performed 100 iterations of the random search by using the above ranges for each of the three prediction experiments. In order to fairly compare the model performances, we matched the number of model parameters to be similar for all baseline methods. To facilitate reproducibility, final hyper-parameter settings we used for all models for each prediction experiments are described at the source code repository, https://github.com/mp2893/gram, along with the detailed steps we used to tune the hyper-parameters.

2 Prediction performance and scalability

Tables 2(a) and 2(b) show the sequential diagnoses prediction performance on Sutter data and MIMIC-III. Both figures show that GRAM+ outperforms other models when predicting labels with significant data insufficiency (i.e. less observed in the training data).The performance gain is greater for MIMIC-III, where GRAM+ outperforms the basic RNN by 10% in the 20th-40th percentile range. This seems to come from the fact that MIMIC patients on average have significantly shorter visit history than Sutter patients, with much more codes received per visit. Such short sequences make it difficult for the RNN to learn and predict diagnoses sequence. The performance difference between GRAM+ and GRAM suggests that our proposed initialization scheme of the basic embeddings ei\mathbf{e}_{i} is important for sequential diagnosis prediction.

Table 2(c) shows the HF prediction performance on Sutter HF cohort. GRAM and GRAM+ consistently outperforms other baselines (except RNN+) by 3∼\sim4% AUC, and RNN+ by maximum 1.8% AUC. These differences are quite significant given that the AUC is already in the mid-80s, a high value for HF prediction, cf. (Choi et al., 2016e). Note that, for GRAM+ and RNN+, we used the downsampled training data to initialize the basic embeddings ei\mathbf{e}_{i}’s and the embedding matrix Wemb\mathbf{W}_{emb} with GloVe, respectively. The result shows that the initialization scheme of the basic embeddings in GRAM+ gives limited improvement over GRAM. This stems from the different natures of the two prediction tasks. While the goal of HF prediction is to predict a binary label for the entire visit sequence, the goal of sequential diagnosis prediction is to predict the co-occurring diagnosis codes at every visit. Therefore the co-occurrence information infused by the initialized embedding scheme is more beneficial to sequential diagnosis prediction. Additionally, this benefit is associated with the natures of the two prediction tasks than the datasets used for the prediction tasks. Because the initialized embedding shows different degrees of improvement as shown by Tables 2(a) and 2(c), when Sutter HF cohort is a subset of Sutter PAMF, thus having similar characteristics.

Overall, GRAM showed superior predictive performance under data insufficiency in three different experiments, demonstrating its general applicability in clinical predictive modeling. Now we briefly discuss the scalability of GRAM by comparing its training time to RNN’s. Table 3 shows the number of seconds taken for the two models to train for a single epoch for each predictive modeling task. GRAM+ and RNN+ showed the similar behavior as GRAM and RNN. GRAM takes approximately 50% more time to train for a single epoch for all prediction tasks. This stems from calculating attention weights and the final representations gi\mathbf{g}_{i} for all medical codes. GRAM also generally takes about 50% more epochs to reach to the model with the lowest validation loss. This is due to optimizing an extra MLP model that generates the attention weights. Overall, use of GRAM adds a manageable amount of overhead in training time to the plain RNN.

3 Qualitative evaluation of interpretable representations

To qualitatively assess the interpretability of the learned representations of the medical codes, we plot on a 2-D space using t-SNE (Maaten and Hinton, 2008) the final representations gi\mathbf{g}_{i} of 2,000 randomly chosen diseases learned by GRAM+ for sequential diagnoses prediction on Sutter dataThe scatterplots of models trained for sequential diagnoses prediction on MIMIC-III and HF prediction for Sutter HF cohort were similar but less structured due to smaller data size. (Figure 3(a)). The color of the dots represents the highest disease categories and the text annotations represent the detailed disease categories in CCS multi-level hierarchy. For comparison, we also show the t-SNE plots on the strongest results from GRAM (Figure 3(b)), RNN+ (Figure 3(c)), RNN (Figure 3(d)) and RandomDAG (Figure 3(e)). GloVe (Figure 3(f)) and Skip-gram (Figure 3(g)) were trained on the Sutter data, where a single visit VtV_{t} was used as the context window to calculate the co-occurrence of codes.

Figures 3(c) and 3(f) confirm that interpretable representations cannot simply be learned only by co-occurrence or supervised prediction without medical knowledge. GRAM+ and GRAM learn interpretable disease

representations that are significantly more consistent with the given knowledge DAG G\mathcal{G}. Based on the prediction performance shown by Table 2, and the fact that the representations gi\mathbf{g}_{i}’s are the final product of GRAM, we can infer that such medically meaningful representations are necessary for predictive models to cope with data insufficiency and make more accurate predictions. Figure 3(b) shows that the quality of the final representations gi\mathbf{g}_{i} of GRAM is quite similar to GRAM+. Compared to other baselines, GRAM demonstrates significantly more structured representations that align well with the given knowledge DAG. It is interesting that Skip-gram shows the most structured representation among all baselines. We used GloVe to initialize the basic embeddings ei\mathbf{e}_{i} in this work because it uses global co-occurrence information and its training time is fast as it is only dependent only on the total number of unique concepts ∣C∣|\mathcal{C}|. Skip-gram’s training time, on the other hand, depends on both the number of patients and the number of visits each patient made, which makes the algorithm generally slower than GloVe. An interactive visualization tool can be accessed at http://www.sunlab.org/research/gram-graph-based-attention-model/.

4 Analysis of the attention behavior

Next we show that GRAM’s attention can be explained intuitively based on the data availability and knowledge DAG’s structure when performing a prediction task. Using Eq. (1), we can calculate the attention weights of individual disease. Figure 4 shows the attention behaviors of four representative diseases when performing HF prediction on Sutter HF cohort.

Other pneumothorax (ICD9 512.89) in Figure 4a is rarely observed in the data and has only five siblings. In this case, most information is derived from the highest ancestor. Temporomandibular joint disorders & articular disc disorder (ICD9 524.63) in Figure 4b is rarely observed but has 139 siblings. In this case, its parent receives a stronger attention because it aggregates sufficient samples from all of its children to learn a more accurate representation. Note that the disease itself also receives a stronger attention to facilitate easier distinction from its large number of siblings.

Unspecified essential hypertension (ICD9 401.9) in Figure 4c is very frequently observed but has only two siblings. In this case, GRAM assigns a very strong attention to the leaf, which is logical because the more you observe a disease, the stronger your confidence becomes. Need for prophylactic vaccination and inoculation against influenza (ICD9 V04.81) in Figure 4d is quite frequently observed and also has 103 siblings. The attention behavior in this case is quite similar to the case with fewer siblings (Figure 4b) with a slight attention shift towards the leaf concept as more observations lead to higher confidence.

Related Work

The attention mechanism is a general framework for neural network learning (Bahdanau et al., 2014), and has been since used in many areas such as speech recognition (Chorowski et al., 2014), computer vision (Ba et al., 2014; Xu et al., 2015) and healthcare (Choi et al., 2016b). However, no one has designed attention model based on knowledge ontology, which is the focus of this work.

There are related works in learning the representations of graphs. Several studies focused on learning the representations of graph vertices by using the neighbor information. DeepWalk (Perozzi et al., 2014) and node2vec (Grover and Leskovec, 2016) use random walk while LINE (Tang et al., 2015) uses breadth-first search to find the neighbors of a vertex and learn its representation based on the neighbor information. Graph convolutional approaches (Yang et al., 2016; Kipf and Welling, 2016) also focus on learning the vertex representations to mainly perform vertex classification. All those works focus on solving the graph data problems whereas GRAM focuses on solving clinical predictive modeling problems using the knowledge DAG as supplementary information.

Several researchers tried to model the knowledge DAG such as WordNet (Miller, 1995) or Freebase (Bollacker et al., 2008) where two entities are connected with various types of relation, forming a set of triples. They aim to project entities and relations (Bordes et al., 2013; Socher et al., 2013; Wang et al., 2014; Lin et al., 2015) to the latent space based on the triples or additional information such as hierarchy of entities (Xie et al., 2016). These works demonstrated tasks such as link prediction, triple classification or entity classification using the learned representations. More recently, Li et al. (2016) learned the representations of words and Wikipedia categories by utilizing the hierarchy of Wikipedia categories. GRAM is fundamentally different from the above studies in that it aims to design intuitive attention mechanism on the knowledge DAG as a knowledge prior to cope with data insufficiency and learn medically interpretable representations to make accurate predictions.

A classical approach for incorporating side information in the predictive models is to use graph Laplacian regularization (Weinberger et al., 2006; Che et al., 2015). However, using this approach is not straightforward as it relies on the appropriate definition of distance on graphs which is often unavailable.

Conclusion

Data insufficiency, either due to less common diseases or small datasets, is one of the key hurdles in healthcare analytics, especially when we apply deep neural networks models. To overcome this challenge, we leverage the knowledge DAG, which provides a multi-resolution view of medical concepts. We propose GRAM, a graph-based attention model using both a knowledge DAG and EHR to learn an accurate and interpretable representations for medical concepts. GRAM chooses a weighted average of ancestors of a medical concept and train the entire process with a predictive model in an end-to-end fashion. We conducted three predictive modeling experiments on real EHR datasets and showed significant improvement in the prediction performance, especially on low-frequency diseases and small datasets. Analysis of the attention behavior provided intuitive insight of GRAM.

References