Graph-Bert: Only Attention is Needed for Learning Graph Representations
Jiawei Zhang, Haopeng Zhang, Congying Xia, Li Sun
Introduction
Graph provides a unified representation for many inter-connected data in the real-world, which can model both the diverse attribute information of the node entities and the extensive connections among these nodes. For instance, the human brain imaging data, online social media and bio-medical molecules can all be represented as graphs, i.e., the brain graph Meng and Zhang 2019, social graph Ugander et al. 2011 and molecular graph Jin et al. 2018, respectively. Traditional machine learning models can hardly be applied to the graph data directly, which usually take the feature vectors as the inputs. Viewed in such a perspective, learning the representations of the graph structured data is an important research task.
In recent years, great efforts have been devoted to designing new graph neural networks (GNNs) for effective graph representation learning. Besides the network embedding models, e.g., node2vec Grover and Leskovec 2016 and deepwalk Perozzi et al. 2014a, the recent graph neural networks, e.g., GCN Kipf and Welling 2016, GAT Veličković et al. 2018 and LoopyNet Zhang 2018, are also becoming much more important, which can further refine the learned representations for specific application tasks. Meanwhile, most of these existing graph representation learning models are still based on the graph structures, i.e., the links among the nodes. Via necessary neighborhood information aggregation or convolutional operators along the links, nodes’ representations learned by such approaches can preserve the graph structure information.
However, several serious learning performance problem, e.g., suspended animation problem Zhang and Meng 2019 and over-smoothing problem Li et al. 2018, with the existing GNN models have also been witnessed in recent years. According to Zhang and Meng 2019, for the GNNs based on the approximated graph convolutional operators Hammond et al. 2011, as the model architecture goes deeper and reaches certain limit, the model will not respond to the training data and suffers from the suspended animation problem. Meanwhile, the node representations obtained by such deep models tend to be over-smoothed and also become indistinguishable Li et al. 2018. Both of these two problems greatly hinder the applications of GNNs for deep graph representation learning tasks. What’s more, the inherently inter-connected nature precludes parallelization within the graph, which becomes critical for large-sized graph input, as memory constraints limit batching across the nodes.
To address the above problems, in this paper, we will propose a new graph neural network model, namely Graph-Bert (Graph based Bert). Inspired by Zhang et al. 2018, model Graph-Bert will be trained with sampled nodes together with their context (which are called linkless subgraphs in this paper) from the input large-sized graph data. Distinct from the existing GNN models, in the representation learning process, Graph-Bert utilizes no links in such sampled batches, which will be purely based on the attention mechanisms instead Vaswani et al. 2017; Devlin et al. 2018. Therefore, Graph-Bert can get rid of the aforementioned learning effectiveness and efficiency problems with existing GNN models promisingly.
What’s more, compared with computer vision He et al. 2018 and natural language processing Devlin et al. 2018, graph neural network pre-training and fine-tuning are still not common practice by this context so far. The main obstacles that prevent such operations can be due to the diverse input graph structures and the extensive connections among the nodes. Also the different learning task objectives also prevents the transfer of GNNs across different tasks. Since Graph-Bert doesn’t really rely on the graph links at all, in this paper, we will investigate the transfer of pre-trained Graph-Bert on new learning tasks and other sequential models (with necessary fine-tuning), which will also help construct the functional pipeline of models in graph learning.
We summarize our contributions of this paper as follows:
New GNN Model: In this paper, we introduce a new GNN model Graph-Bert for graph data representation learning. Graph-Bert doesn’t rely on the graph links for representation learning and can effectively address the suspended animation problems aforementioned. Also Graph-Bert is trainable with sampled linkless subgraphs (i.e., target node with context), which is more efficient than existing GNNs constructed for the complete input graph. To be more precise, the training cost of Graph-Bert is only decided by (1) training instance number, and (2) sampled subgraph size, which is uncorrelated with the input graph size at all.
Unsupervised Pre-Training: Given the input unlabeled graph, we will pre-train Graph-Bert based on to two common tasks in graph studies, i.e., node attribute reconstruction and graph structure recovery. Node attribute recovery ensures the learned node representations can capture the input attribute information; whereas graph structure recovery can further ensure Graph-Bert learned with linkless subgraphs can still maintain both the graph local and global structure properties.
Fine-Tuning and Transfer: Depending on the specific application task objectives, the Graph-Bert model can be further fine-tuned to adapt the learned representations to specific application requirements, e.g., node classification and graph clustering. Meanwhile, the pre-trained Graph-Bert can also be transferred and applied to other sequential models, which allows the construction of functional pipelines for graph learning.
The remaining parts of this paper are organized as follows. We will introduce the related work in Section 2. Detailed information about the Graph-Bert model will be introduced in Section 3, whereas the pre-training and fine-tuning of Graph-Bert will be introduced in Section 4 in detail. The effectiveness of Graph-Bert will be tested in Section 5. Finally, we will conclude this paper in Section 6.
Related Work
To make this paper self-contained, we will introduce some related topics here on GNNs, Transformer and Bert.
Graph Neural Network: Representative examples of GNNs proposed by present include GCN Kipf and Welling 2016, GraphSAGE Hamilton et al. 2017 and LoopyNet Zhang 2018, based on which various extended models Veličković et al. 2018; Sun et al. 2019; Klicpera et al. 2018 have been introduced as well. As mentioned above, GCN and its variant models are all based on the approximated graph convolutional operator Hammond et al. 2011, which may lead to the suspended animation problem Zhang and Meng 2019 and over-smoothing problem Li et al. 2018 for deep model architectures. Theoretic analyses of the reasons are provided in Li et al. 2018; Zhang and Meng 2019; Gürel et al. 2019. To handle such problems, Zhang and Meng 2019 generalizes the graph raw residual terms in Zhang 2018 and proposes a method based on graph residual learning; Li et al. 2018 proposes to adopt residual/dense connections and dilated convolutions into the GCN architecture. Several other works Sun et al. 2019; Huang and Carley 2019 seek to involve the recurrent network for deep graph representation learning instead.
Bert and Transformer: In NLP, the dominant sequence transduction models are based on complex recurrent Hochreiter and Schmidhuber 1997; Chung et al. 2014 or convolutional neural networks Kim 2014. However, the inherently sequential nature precludes parallelization within training examples. Therefore, in Vaswani et al. 2017, the authors propose a new network architecture, the Transformer, based solely on attention mechanisms, dispensing with recurrence and convolutions entirely. With Transformer, Devlin et al. 2018 further introduces Bert for deep language understanding, which obtains new state-of-the-art results on eleven natural language processing tasks. In recent years, Transformer and Bert based learning approaches have been used extensively in various learning tasks Dai et al. 2019; Lan et al. 2019; Shang et al. 2019.
Readers may also refer to page https://paperswithcode.com/area/graphs and page https://paperswithcode.com/area/natural-language-processing for more information on the state-of-the-art work on these topics.
Method
In this section, we will introduce the detailed information about the Graph-Bert model. As illustrated in Figure 1, Graph-Bert involves several parts: (1) linkless subgraph batching, (2) node input embedding, (3) graph-transformer based encoder, (4) representation fusion, and (5) the functional component. The results learned by the the graph-transformer model will be fused as the representation for the target nodes. In this section, we will introduce these key parts in great detail, whereas the pre-training and fine-tuning of Graph-Bert will be introduced in the following section.
In the sequel of this paper, we will use the lower case letters (e.g., ) to represent scalars, lower case bold letters (e.g., ) to denote column vectors, bold-face upper case letters (e.g., ) to denote matrices, and upper case calligraphic letters (e.g., ) to denote sets or high-order tensors. Given a matrix , we denote and as its row and column, respectively. The (, ) entry of matrix can be denoted as either . We use and to represent the transpose of matrix and vector . For vector , we represent its -norm as . The Frobenius-norm of matrix is represented as . The element-wise product of vectors and of the same dimension is represented as , whose concatenation is represented as .
2 Linkless Subgraph Batching
There exist different metrics to measure the intimacy scores among the nodes within the graph, e.g., Jaccard’s coefficienty Jaccard 1901, Adamic/Adar Adamic and Adar 2003, Katz Katz 1953. In this paper, we define matrix based on the pagerank algorithm, which can be denoted as
where factor (which is usually set as ). Term denotes the colum-normalized adjacency matrix. In its representation, is the adjacency matrix of the input graph, and is its corresponding diagonal matrix with on its diagonal.
Formally, for any target node in the input graph, based on the intimacy matrix , we can define its learning context as follows:
(Node Context): Given an input graph and its intimacy matrix , for node in the graph, we define its learning context as set . Here, the term defines the minimum intimacy score threshold for nodes to involve in ’s context.
We may need to add a remark: for all the nodes in ’ learning context , they can cover both local neighbors of as well as the nodes which are far away. In this paper, we define the threshold as the entry of (with being excluded), i.e., covers the top-k intimate nodes of in graph . Based on the node context concept, we can also represent the set of sampled graph batches for all the nodes as set , and denotes the subgraph sampled for (as the target node). Formally, can be represented as , where the node set covers both and its context nodes and the link set is null. For large-sized input graphs, set can further be decomposed into several mini-batches, i.e., , which will be fed to train the Graph-Bert model.
3 Node Input Vector Embeddings
Different from image and text data, where the pixels and words/chars have their inherent orders, nodes in graphs are orderless. The Graph-Bert model to be learned in this paper doesn’t require any node orders of the input sampled subgraph actually. Meanwhile, to simplify the presentations, we still propose to serialize the input subgraph nodes into certain ordered list instead. Formally, for all the nodes in the sampled linkless subgraph , we can denote them as a node list , where will be placed ahead of if . For the remaining of this subsection, we will follow the identical node orders as indicated above by default to define their input vector embeddings.
The input vector embeddings to be fed to the graph-transformer model actually cover four parts: (1) raw feature vector embedding, (2) Weisfeiler-Lehman absolute role embedding, (3) intimacy based relative positional embedding, and (4) hop based relative distance embedding, respectively.
Formally, for each node in the subgraph , we can embed its raw feature vector into a shared feature space (of the same dimension ) with its raw feature vector , which can be denoted as
Depending on the input raw features properties, different models can be used to define the function. For instance, CNN can be used if denotes images; LSTM/BERT can be applied if denotes texts; and simple fully connected layers can also be used for simple attribute inputs.
3.2 Weisfeiler-Lehman Absolute Role Embedding
3.3 Intimacy based Relative Positional Embedding
The WL based role embeddings can capture the global node role information in the representations. Here, we will introduce a relative positional embedding to extract the local information in the subgraph based on the placement orders of the serialized node list introduced at the beginning of this subsection. Formally, based on that serialized node list, we can denote the position of as . We know that by default and nodes closer to will have a small positional index. Furthermore, is a variant position index metric. For the identical node , its positional index will be different for different sampled subgraphs.
Formally, for node , we can also extract its intimacy based relative positional embedding with the function defined above as follows:
which is quite close to the positional embedding in Vaswani et al. 2017 for the relative positions in the word sequence.
3.4 Hop based Relative Distance Embedding
The hop based relative distance embedding can be treated as a balance between the absolute role embedding (for global information) and intimacy based relative positional embedding (for local information). Formally, for node in the subgraph , we can denote its relative distance in hops to in the original input graph as , which can be used to define its embedding vector as
It it easy to observe that vector will also be variant for the identical node in different subgraphs.
4 Graph Transformer based Encoder
Based on the computed embedding vectors defined above, we will be able to aggregate them together to define the initial input vectors for nodes, e.g., , in the subgraph as follows:
Graph-Bert Learning
We propose to pre-train Graph-Bert with two tasks: (1) node attribute reconstruction, and (2) graph structure recovery. Meanwhile, depending on the objective application tasks, e.g., (1) node classification and (2) graph clustering as studied in this paper, Graph-Bert can be further fine-tuned to adapt both the model and the learned node representations accordingly to the new tasks.
The node raw attribute reconstruction task focuses on capturing the node attribute information in the learned representations, whereas the graph structure recovery task focuses more on the graph connection information instead.
Formally, for the target node in the sampled subgraph , we have its learned representation by Graph-Bert to be . Via the fully connected layer (together with the activation function layer if necessary), we can denote the reconstructed raw attributes for node based on as . To ensure the learned representations can capture the node raw attribute information, compared against the node raw features, e.g., for , we can define the node raw attribute reconstruction based loss term as follows:
1.2 Task #2: Graph Structure Recovery
Furthermore, to ensure such representation vectors can also capture the graph structure information, the graph structure recovery task is also used as a pre-training task. Formally, for any two nodes and , based on their learned representations, we can denote the inferred connection score between them by computing their cosine similarity, i.e., . Compared against the ground truth graph intimacy matrix defined in Section 3.2, i.e., , we can denote the introduced loss term as follows:
2 Model Transfer and Fine-tuning
In applying the learned Graph-Bert into new learning tasks, the learned graph representations can be either fed into the new tasks directly or with necessary adjustment, i.e., fine-tuning. In this part, we can take the node classification and graph clustering tasks as the examples, where graph clustering can use the learned representations directly but fine-tuning will be necessary for the node classification task.
Based on the nodes learned representations, e.g., for , we can denote the inferred label for the node via the functional component as . Compared with the nodes’ true labels, we will be able to define the introduced node classification loss term on training batch as
By re-training these stacked fully connected layers together with Graph-Bert (loaded from pre-training), we will be able to infer node class labels.
2.2 Task # 2: Graph Clustering
The above objective function involves multiple variables to be learned concurrently, which can be trained with the EM algorithm much more effectively instead of error backpropagation. Therefore, instead of re-training the above graph clustering model together with Graph-Bert, we will only take the learned node representations as the node feature input for learning the graph clustering model instead.
Experiments
To test the effectiveness of Graph-Bert in learning the graph representations, in this section, we will provide extensive experimental results of Graph-Bert on three real-world benchmark graph datasets, i.e., Cora, Citeseer and Pubmed Yang et al. 2016, respectively.
Reproducibility. Both the datasets and source code used can be accessed via link https://github.com/jwzhanggy/Graph-Bert. Detailed information about the server used to run the model can be found at the footnote GPU Server: ASUS X99-E WS motherboard, Intel Core i7 CPU 6850K@3.6GHz (6 cores), 3 Nvidia GeForce GTX 1080 Ti GPU (11 GB buffer each), 128 GB DDR4 memory and 128 GB SSD swap..
The graph benchmark datasets used in the experiments include Cora, Citeseer and Pubmed Yang et al. 2016, which are used in most of the recent state-of-the-art graph neural network research works Kipf and Welling 2016; Veličković et al. 2018; Zhang and Meng 2019. Based on the input graph data, we will first pre-compute the node intimacy scores, based on which subgraph batches will be sampled subject to the subgraph size . In addition, we will also pre-compute the node pairwise hop distance and WL node codes. By minimizing the node raw feature reconstruction loss and graph structure recovery loss, Graph-Bert can be effectively pre-trained, whose learned variables will be transferred to the follow-up node classification and graph clustering tasks with/without fine-tuning. In the experiments, we first pre-train Graph-Bert based on the node attribute reconstruction task with 200 epochs, then load and pre-train the same Graph-Bert model again based on the graph structure recovery task with another 200 epochs. In Figure 2, we show the learning performance of Graph-Bert on node attribute reconstruction and graph recovery, which converges very fast on both of these tasks.
If not clearly specified, the results reported in this paper are based on the following parameter settings of Graph-Bert: subgraph size: (Cora), (Citeseer) and (Pubmed); hidden size: 32; attention head number: 2; hidden layer number: ; learning rate: 0.01 (Cora) and 0.001 (Citeseer) and 0.0005 (Pubmed); weight decay: ; intermediate size: 32; hidden dropout rate: 0.5; attention dropout rate: 0.3; graph residual term: graph-raw; training epoch: 150 (Cora), 500 (Pubmed), 2000 (Citeseer).
2 Node Classification without Pre-training
Graph-Bert is a powerful mode and it can be applied to address various graph learning tasks in the standalone mode. To show the effectiveness of Graph-Bert, we will first provide the experimental results of Graph-Bert on the node classification task without pre-training here, whereas the pre-trained Graph-Bert based node classification results will be provided in Section 5.4 in more detail. Here, we will follow the identical train/validation/test set partitions used in the existing graph neural network papers Yang et al. 2016 for fair comparisons.
In Figure 3, we illustrate the training records of Graph-Bert for node classification on the Cora dataset. To show that Graph-Bert is different from other GNN models and Graph-Bert works with deep architectures, we also change the model depth with values from . According to the plots, Graph-Bert can converge very fast (with less than 10 epochs) on the training set. What’s more, as the model depth increases, Graph-Bert will not suffer from the suspended animation problem. Even the very deep Graph-Bert (50 layers) can still respond effectively to the training data and achieve good learning performance.
2.2 Main Results
The learning results of Graph-Bert (with different graph residual terms) on node classification are provided in Table 1. The comparison methods used here cover both classic and state-of-the-art GNN models. For the variant models which extend GCN and GAT (with new learning settings, include more training data, re-configure the graph structure or use new optimization methods), we didn’t compare them here. However, similar techniques proposed by these extension works can be used to further help improve Graph-Bert as well. According to the achieved scores, we observe that Graph-Bert can out-perform most of these baseline methods with a big improvement on both Cora and Pubmed. On Citeseer, its perofrmance is also among the top 3.
2.3 Subgraph Size kk Analysis
As illustrated in Table 2, we provide the learning performance analysis of Graph-Bert with different subgraph sizes, i.e., parameter , on the Cora dataset. According to the results, parameter affects the learning performance of Graph-Bert a lot, since it defines how many nearby nodes will be used to define the nodes’ learning context. For the Cora dataset, we observe that the learning performance of Graph-Bert improves steadily as increases from to . After that, as further increases, the performance will degrade dramatically. For the good scores with , partial contributions come from the graph residual terms in Graph-Bert. The time cost of Graph-Bert increases as goes larger, which is very minor actually compared with other existing GNN models, like GCN and GAT. Similar results can be observed for the other two datasets, but the optimal are different.
2.4 Graph Residual Analysis
What’s more, in Table 3, we also provide the learning results of Graph-Bert with different graph residual terms. According to the scores, Graph-Bert with graph-raw residual term can outperform the other two, which is also consistent with the experimental observations on these different residual terms as reported in Zhang and Meng 2019.
2.5 Initial Embedding Analysis
As shown in Table 4, we provide the learning performance of Graph-Bert on these three datasets, which takes different initial embeddings as the input. To better show the performance differences, the Graph-Bert used here doesn’t involve any residual learning. According to the results, using the Weisfeiler-Lehman role embedding, hop based distance embedding and intimacy based positional embedding vectors along, Graph-Bert cannot work very well actually, whereas the raw feature embeddings do contribute a lot. Meanwhile, by incorporating such complementary embeddings into the raw feature embedding, the model can achieve better performance than using raw feature embedding only.
3 Graph Clustering without Pre-Training
In Table 5, we show the learning results of Graph-Bert on graph clustering without any pre-training on the three datasets. Formally, the clustering component used in Graph-Bert is KMeans, which takes the nodes’ raw feature vectors as the input. The results are evaluated with several different metrics shown above.
4 Pre-training vs. No Pre-training
The results reported in the previous subsections are all based on the Graph-Bert without pre-training actually. Here, we will provide the experimental results on Graph-Bert with pre-training to show their differences. According to the experiments, given enough training epochs, models with/without pre-training can both converge to very good learning results. Therefore, to highlight the differences, we will only use of the normal training epochs here for fine-tuning Graph-Bert, and the results are provided in Table 6. We also show the performance of Graph-Bert without pre-training here for comparison.
According to the scores, for most of the datasets, pre-training do give Graph-Bert a good initial state, which helps the model achieve better performance with only a very small number of fine-tuning epochs. On Cora and Citeseer, pre-training helps both the node classification and graph clustering tasks. Meanwhile, for Pubmed, pre-training helps node classification but degrades the results on graph clustering. Also pre-training with both node classification and graph recovery help the model to capture more information from the graph data, which also lead to higher scores than the models with single pre-training tasks.
Conclusion
In this paper, we have introduced the new Graph-Bert model for graph representation learning. Different from existing GNNs, Graph-Bert works well in deep architectures and will not suffer from the common problems with other GNNs. Based on a batch of linkless subgraphs sampled from the original graph data, Graph-Bert can effectively learn the representations of the target node with the extended graph-transformer layers introduced in this paper. Graph-Bert can serve as the graph representation learning component in graph learning pipeline. The pre-trained Graph-Bert can be transferred and applied to address new tasks either directly or with necessary fine-tuning.