Watch Your Step: Learning Node Embeddings via Graph Attention
Sami Abu-El-Haija, Bryan Perozzi, Rami Al-Rfou, Alex Alemi
Introduction
Unsupervised graph embedding methods seek to learn representations that encode the graph structure. These embeddings have demonstrated outstanding performance on a number of tasks including node classification , knowledge-base completion , semi-supervised learning , and link prediction . In general, as introduced by Perozzi et al , these methods operate in two discrete steps: First, they sample pair-wise relationships from the graph through random walks and counting node co-occurances. Second, they train an embedding model e.g. using Skipgram of word2vec , to learn representations that encode pairwise node similarities.
While such methods have demonstrated positive results on a number of tasks, their performance can significantly vary based on the setting of their hyper-parameters. For example, observed that the quality of learned representations is dependent on the length of the random walk . In practice, DeepWalk and many of its extensions [e.g. 13] use word2vec implementations . Accordingly, it has been revealed by that the hyper-parameter , refered to as training window length in word2vec , actually controls more than a fixed length of the random walk. Instead, it parameterizes a function, we term the context distribution and denote , which controls the probability of sampling a node-pair when visited within a specific distance Specifically, rather than using as constant and assuming all nodes visited within distance are related, a desired context distance is sampled from uniform () for each node pair in training. If the node pair was visited more than -steps apart, it is not used for training. This was revealed to us by Levy et al , see their Section 3.1. Many DeepWalk-style methods inherited this, as they utilize word2vec implementation.. Implicitly, the choices of and , create a weight mass on every node’s neighborhood. In general, the weight is higher on nearby nodes, but the specific form of the of the mass function is determined by the aforementioned hyper-parameters. In this work, we aim to replace these hyper-parameters with trainable parameters, so that they can be automatically learned for each graph. To do so, we pose graph embedding as end-to-end learning, where the (discrete) two steps of random walk co-occurance sampling, followed by representation learning, are joint using a closed-form expectation over the graph adjacency matrix.
Our inspiration comes from the successful application of attention models in domains such as Natural Language Processing (NLP) [e.g. 4, 36], image recognition , and detecting rare events in videos . To the best of our knowledge, the approach we propose is significantly different from the standard application of attention models. Instead of using attention parameters to guide the model where to look when making a prediction, we use attention parameters to guide our learning algorithm to focus on parts of the data that are most helpful for optimizing upstream objective.
We show mathematical equivalence between the context distribution and the co-efficients of power series of the transition matrix. This allows us to learn the context distribution by learning an attention model on the power series. The attention parameters “guide” the random walk, by allowing it to focus more on short- or long-term dependencies, as best suited for the graph, while optimizing an upstream objective. To the best of our knowledge, this work is the first application of attention methods to graph embedding.
Specifically, our contributions are the following:
We propose an extendible family of graph attention models that can learn arbitrary (e.g. non-monotonic) context distributions.
We show that the optimal choice of context distribution hyper-parameters for competing methods, found by manual tuning, agrees with our automatically-found attention parameters.
We evaluate on a number of challenging link prediction tasks comprised of real world datasets, including social, collaboration, and biological networks. Experiments show we substantially improve on our baselines, reducing link-prediction error by 20%-40%.
Preliminaries
2 Learning Embeddings via Random Walks
Introduced by , this family of methods [incl. 13, 18] induce random walks along by starting from a random node , and repeatedly sampling an edge to transition to next node as , where are the outgoing edges from . The transition sequences (i.e. random walks) can then be passed to word2vec algorithm, which learns embeddings by stochastically taking every node along the sequence , and the embedding representation of this anchor node is brought closer to the embeddings of its next neighbors, , the context nodes. In practice, the context window size is sampled from a distribution e.g. uniform as explained in .
where partition function can be estimated with negative sampling .
A recently-proposed objective for learning embeddings is the graph likelihood :
where is the output of the model evaluated at edge , given node embedings ; the activation function is the logistic; Maximizing the graph likelihood pushes the model score towards if value is large and pushes it towards if .
In our work, we minimize the negative log of Equation 3, written in our matrix notation as:
3 Attention Models
We mention attention models that are most similar to ours [e.g. 25, 29, 33] , where an attention function is employed to suggest positions within the input example that the classification function should pay attention to, when making inference. This function is used during the training phase in the forward pass and in the testing phase for prediction. The attention function and the classifier are jointly trained on an upstream objective e.g. cross entropy. In our case, the attention mechanism is only guides the learning procedure, and not used by the model for inference. Our mechanism suggests parts of the data to focus on, during training, as explained next.
Our Method
Let be the transition matrix for a graph, which can be calculated by normalizing the rows of to sum to one. This can be written as:
Eq. (7) is derived, step-by-step, in the Appendix. We are not concerned by the exact definition of the scalar coefficient, , but we note that the coefficient decreases with .
Instead of keeping a hyper-parameter, we want to analytically optimize it on an upstream objective. Further, we are interested to learn the co-efficients to instead of hand-engineering a formula.
2 Learning the Context Distribution
We want to learn the co-efficients to . Let the context distribution be a -dimensional vector as with and . We assign co-efficient to . Formally, our expectation on is parameterized with, and is differentiable w.r.t., :
Training embeddings over random walk sequences, using word2vec or GloVe, respectively, are special cases of Equation 8, with fixed apriori as or .
3 Graph Attention Models
To learn automatically, we propose an attention model which guides the random surfer on “where to attend to” as a function of distance from the source node. Specifically, we define a Graph Attention Model as a process which models a node’s context distribution as the output of softmax:
where the variables are trained via backpropagation, jointly while learning node embeddings. Our hypothesis is as follows. If we don’t impose a specific formula on , other than (regularized) softmax, then we can use very large values of and allow every graph to learn its own form of with its preferred sparsity and own decay form. Should the graph structure require a small , then the optimization would discover a left-skewed with all of probability mass on and . However, if according to the objective, a graph is more accurately encoded by making longer walks, then they can learn to use a large (e.g. using uniform or even right-skewed Q distribution), focusing more attention on longer distance connections in the random walk.
To this end, we propose to train softmax attention model on the infinite power series of the transition matrix. We define an expectation on our proposed random walk matrix asWe do not actually unroll the summation in Eq. (10) an infinite number of times. Our experiments show that unrolling it 10 or 20 times is sufficient to obtain state-of-the-art results. :
where are jointly trained with the embeddings to minimize our objective.
4 Training Objective
The final training objective for the Softmax attention mechanism, coming from the NLGL Eq. (4),
5 Algorithmic Complexity
The naive computation of requires matrix multiplications and so is . However, as most real-world adjacency matrices have an inherent low rank structure, a number of fast approximations to computing the random walk transition matrix raised to a power have been proposed [e.g. 32]. Alternatively SVD can decompose as and then the power can be calculated by raising the diagonal matrix of singular values to as since . Furthermore, the SVD can be approximated in time linear to the number of non-zero entries . Therefore, we can calculate in .
6 Extensions
As presented, our proposed method can learn the weights of the context distribution . However, we briefly note that such a model can be trivially extended to learn the weight of any other type of pair-wise node similarity (e.g. Personalized PageRank, Adamic-Adar, etc). In order to do this, we can extend the definition of the context with an additional dimension for the new type of similarity, and an additional element in the softmax to learn a joint importance function.
Experiments
We evaluate the quality of embeddings produced when random walks are augmented with attention, through experiments on link prediction . Link prediction is a challenging task, with many real world applications in information retrieval, recommendation systems and social networks. As such, it has been used to study the properties of graph embeddings . Such an intrinsic evaluation emphasizes the structure-preserving properties of embedding.
Our experimental setup is designed to determine how well the embeddings produced by a method captures the topology of the graph. We measure this in the manner of : remove a fraction (=50%) of graph edges, learn embeddings from the remaining edges, and measure how well the embeddings can recover those edges which have been removed. More formally, we split the graph edges into two partitions of equal size and such that the training graph is connected. We also sample non existent edges () to make and . We use (, ) for training and model selection, and use (, ) to compute evaluation metrics.
Datasets: Table 1(a) describes the datasets used in our experiments. Datasets available from SNAP https://snap.stanford.edu/data.
Results: Our results, summarized in Table 1, show that our proposed methods substantially outperform all baseline methods. Specifically, we see that the error is reduced by up to over baseline methods which have fixed context definitions. This shows that by parameterizing the context distribution and allowing each graph to learn its own distribution, we can better preserve the graph structure (and thereby better predict missing edges).
Discussion: Figure 2(a) shows how the learned attention weights vary across datasets. Each dataset learns its own attention form, and the highest weights generally correspond to the highest weights when doing a grid search over for node2vec (as in Figure 1(b)).
The hyper-parameter determines the highest power of the transition matrix, and hence the maximum context size available to the attention model. We suggest using large values for , since the attention weights can effectively use a subset of the transition matrix powers. For example, if a network needs only 2 hops to be accurately represented, then it is possible for the softmax attention model to learn . Figure 2(b) shows how varying the regularization term allows the softmax attention model to “attend to” only what each dataset requires. We observe that for most graphs, the majority of the mass gets assigned to . This shows that shorter walks are more beneficial for most graphs. However, on wiki-vote, better embeddings are produced by paying attention to longer walks, as its softmax is uniform-like, with a slight right-skew.
2 Sensitivity Analysis
So far, we have removed two hyper-parameters, the maximum window size , and the form of the context distribution . In exchange, we have introduced other hyper-parameters – specifically walk length (also ) and a regularization term for the softmax attention model. Nonetheless, we show that our method is robust to various choices of these two. Figures 2(a) and 2(b) both show that the softmax attention weights drop to almost zero if the graph can be preserved using shorter walks, which is not possible with fixed-form distributions (e.g. ).
Figure 4 examines this relationship in more detail for dimensional embeddings, sweeping our hyper-parameters and , and comparing results to the best and worst node2vec embeddings for . (Note that node2vec lines are horizontal, as they do not depend on .) We observe that all the accuracy metrics are within to , when varying these hyper-parameters, and are all still well-above our baseline (which sample from a fixed-form context distribution).
3 Node Classification Experiments
Our classification prediciton function contains one scalar parameter . It can be thought of a “smooth” k-nearest-neighbors, as it takes a weighted average of known labels, where the weights are exponential of the dot-product similarity. Such a simple function should introduce no model bias.
Related Work
The field of learning on graphs has attracted much attention lately. Here we summarize two broad classes of algorithms, and point the reader to several recent reviews for more context.
The first class of algorithms are semi-supervised and concerned with predicting labels over a graph, its edges, and/or its nodes. Typically, these algorithms process a graph (nodes and edges) as well as per-node features. These include recent graph convolution methods [e.g. 26, 7, 3, 15, 33] with spectral variants , diffusion methods [e.g. 10, 12, 9], including ones trained until fixed-point convergence and semi-supervised node classification with low-rank approximation of convolution . We differ from these methods as (1) our algorithm is unsupervised (trained exclusively from the graph structure itself) without utilizing labels during training, and (2) we explicitly model the relationship between all node pairs.
Conclusion
In this paper, we propose an attention mechanism for learning the context distribution used in graph embedding methods. We derive the closed-form expectation of the DeepWalk co-occurrence statistics, showing an equivalence between the context distribution hyper-parameters, and the co-efficients of the power series of the graph transition matrix. Then, we propose to replace the context hyper-parameters with trainable models, that we learn jointly with the embeddings on an objective that preserves the graph structure (the Negative Log Graph Likelihood, NLGL). Specifically, we propose Graph Attention Models, using a softmax to learn a free-form contexts distribution with a parameter for each type of context similarity (e.g. distance in a random walk).
We show significant improvements on link prediction and node classification over state-of-the-art baselines (that use a fixed-form context distribution), reducing error on link prediction and classification, respectively by up to 40% and 10%. In addition to improved performance (by learning distributions of arbitrary forms), our method can obviate the manual grid search over hyper-parameters: walk length and form of context distribution, which can drastically fluctuate the quality of the learned embeddings and are different for every graph. On the datasets we consider, we show that our method is robust to its hyper-parameters, as described in Section 4.2. Our visualizations of converged attention weights convey to us that some graphs (e.g. voting graphs) can be better preserved by using longer walks, while other graphs (e.g. protein-protein interaction graphs) contain more information in short dependencies and require shorter walks.
We believe that our contribution in replacing these sampling hyperparameters with a learnable context distribution is general and can be applied to many domains and modeling techniques in graph representation learning.
References
Appendix
Now, if node was visited steps after node , then the probabilitiy of it being sampled is given by:
In case of DeepWalk , probability above equals:
and event is independant of the condition . Further, event can be partitioned and Eq. (14) can be written as
where second line is trivial since the events are disjoint. We can now use Bayes’ rule to derive the probability of being visited steps after and being selected in ’s sampled context, as:
Now, let be the event that a walker visits and after steps, visits and selects it part of its context. This event happens with the probability indicated in Equation 18. Concretely,
Let count the events , then:
Suppose we run DeepWalk, starting random walks from each node , then the expected number of times that is present in the context of is given by:
Finally, we can write down the expectation over the square matrix :
3 Depiction of Learned Context Distribution
& Figure: Depiction of how our model assigns context distributions (shaded red) compared to earlier work. We depict the graph from the perspective of anchor node (yellow). Given a social graph (top), where friends of friends are usually friends, our algorithm learns a leftskewed distribution. Given a voting graph (bottom), with general transitivity: , it learns a long-tail distribution. Earlier methods (e.g. DeepWalk) use word2vec, which internally uses a linear decay context distribution, treating all graphs the same.