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 (C)(C). In practice, DeepWalk and many of its extensions [e.g. 13] use word2vec implementations . Accordingly, it has been revealed by that the hyper-parameter CC, 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 QQ, which controls the probability of sampling a node-pair when visited within a specific distance Specifically, rather than using CC as constant and assuming all nodes visited within distance CC are related, a desired context distance cic_{i} is sampled from uniform (ci∼U{1,C}c_{i}\sim\mathcal{U}\{1,C\}) for each node pair ii in training. If the node pair ii was visited more than cic_{i}-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 CC and QQ, 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 EE by starting from a random node v0∈sample(V)v_{0}\in\textit{sample}(V), and repeatedly sampling an edge to transition to next node as vi+1:=sample(E[vi])v_{i+1}:=\textit{sample}(E[v_{i}]), where E[vi]E[v_{i}] are the outgoing edges from viv_{i}. The transition sequences v0→v1→v2→…v_{0}\rightarrow v_{1}\rightarrow v_{2}\rightarrow\dots (i.e. random walks) can then be passed to word2vec algorithm, which learns embeddings by stochastically taking every node along the sequence viv_{i}, and the embedding representation of this anchor node viv_{i} is brought closer to the embeddings of its next neighbors, {vi+1,vi+2,…,vi+c}\{v_{i+1},v_{i+2},\dots,v_{i+c}\}, the context nodes. In practice, the context window size cc is sampled from a distribution e.g. uniform U{1,C}\mathcal{U}\{1,C\} as explained in .

where partition function Z=∑v,uexp⁡(Yv⊤Yu)Z=\sum_{v,u}\exp(Y_{v}^{\top}Y_{u}) can be estimated with negative sampling .

A recently-proposed objective for learning embeddings is the graph likelihood :

where g(Y)v,ug(\mathbf{Y})_{v,u} is the output of the model evaluated at edge (v,u)(v,u), given node embedings Y\mathbf{Y}; the activation function σ(.)\sigma(.) is the logistic; Maximizing the graph likelihood pushes the model score g(Y)v,ug(\mathbf{Y})_{v,u} towards 11 if value DvuD_{vu} is large and pushes it towards if (v,u)∉E(v,u)\notin E.

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 T\mathcal{T} be the transition matrix for a graph, which can be calculated by normalizing the rows of A\mathbf{A} 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, [1−k−1C]\left[1-\frac{k-1}{C}\right], but we note that the coefficient decreases with kk.

Instead of keeping CC a hyper-parameter, we want to analytically optimize it on an upstream objective. Further, we are interested to learn the co-efficients to (T)k\left(\mathcal{T}\right)^{k} instead of hand-engineering a formula.

2 Learning the Context Distribution

We want to learn the co-efficients to (T)k\left(\mathcal{T}\right)^{k}. Let the context distribution QQ be a CC-dimensional vector as Q=(Q1,Q2,⋯ ,QC)Q=(Q_{1},Q_{2},\cdots,Q_{C}) with Qk≥0Q_{k}\geq 0 and ∑kQk=1\sum_{k}Q_{k}=1. We assign co-efficient QkQ_{k} to (T)k\left(\mathcal{T}\right)^{k}. Formally, our expectation on D\mathbf{D} is parameterized with, and is differentiable w.r.t., QQ:

Training embeddings over random walk sequences, using word2vec or GloVe, respectively, are special cases of Equation 8, with QQ fixed apriori as Qk=[1−k−1C]Q_{k}=\left[1-\frac{k-1}{C}\right] or Qk∝1kQ_{k}\propto\frac{1}{k}.

3 Graph Attention Models

To learn QQ 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 QQ as the output of softmax:

where the variables qkq_{k} are trained via backpropagation, jointly while learning node embeddings. Our hypothesis is as follows. If we don’t impose a specific formula on Q=(Q1,Q2,…QC)Q=(Q_{1},Q_{2},\dots Q_{C}), other than (regularized) softmax, then we can use very large values of CC and allow every graph to learn its own form of QQ with its preferred sparsity and own decay form. Should the graph structure require a small CC, then the optimization would discover a left-skewed QQ with all of probability mass on {Q1,Q2}\{Q_{1},Q_{2}\} and ∑k>2Qk≈0\sum_{k>2}Q_{k}\approx 0. However, if according to the objective, a graph is more accurately encoded by making longer walks, then they can learn to use a large CC (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 Dsoftmax[∞]\mathbf{D}^{\text{softmax}[\infty]} 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 q1,q2,…q_{1},q_{2},\dots 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 (T)k(\mathcal{T})^{k} requires kk matrix multiplications and so is O(∣V∣3k)\mathcal{O}(|V|^{3}k). 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 kk have been proposed [e.g. 32]. Alternatively SVD can decompose T\mathcal{T} as T=UΛVT\mathcal{T}=\mathcal{U}\Lambda\mathcal{V}^{T} and then the kthk^{\textrm{th}} power can be calculated by raising the diagonal matrix of singular values to kk as (T)k=U(Λ)kVT(\mathcal{T})^{k}=\mathcal{U}(\Lambda)^{k}\mathcal{V}^{T} since VTU=I\mathcal{V}^{T}\mathcal{U}=I. Furthermore, the SVD can be approximated in time linear to the number of non-zero entries . Therefore, we can calculate (T)k(\mathcal{T})^{k} in O(∣E∣)\mathcal{O}(|E|).

6 Extensions

As presented, our proposed method can learn the weights of the context distribution C\mathcal{C}. 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 QQ with an additional dimension Qk+1Q_{k+1} for the new type of similarity, and an additional element in the softmax qk+1q_{k+1} 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 EE into two partitions of equal size EtrainE_{\text{train}} and EtestE_{\text{test}} such that the training graph is connected. We also sample non existent edges ((u,v)∉E(u,v)\notin E) to make Etrain−E^{-}_{\text{train}} and Etest−E_{\text{test}}^{-}. We use (EtrainE_{\text{train}}, Etrain−E^{-}_{\text{train}}) for training and model selection, and use (EtestE_{\text{test}}, Etest−E_{\text{test}}^{-}) 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 45%45\% 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 QQ 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 CC for node2vec (as in Figure 1(b)).

The hyper-parameter CC 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 CC, 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 Q3,Q4,⋯≈0Q_{3},Q_{4},\dots\approx 0. Figure 2(b) shows how varying the regularization term β\beta 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 Q1,Q2Q_{1},Q_{2}. 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 QQ is uniform-like, with a slight right-skew.

2 Sensitivity Analysis

So far, we have removed two hyper-parameters, the maximum window size CC, and the form of the context distribution U\mathcal{U}. In exchange, we have introduced other hyper-parameters – specifically walk length (also CC) and a regularization term β\beta 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. U\mathcal{U}).

Figure 4 examines this relationship in more detail for d=128d=128 dimensional embeddings, sweeping our hyper-parameters CC and β\beta, and comparing results to the best and worst node2vec embeddings for C∈C\in. (Note that node2vec lines are horizontal, as they do not depend on β\beta.) We observe that all the accuracy metrics are within 1%1\% to 2%2\%, 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 α\alpha. 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 uu was visited kk steps after node vv, then the probabilitiy of it being sampled is given by:

In case of DeepWalk , probability above equals:

and event k≤ck\leq c is independant of the condition (x(k)=u∩x(0)=v)\left(x(k)=u\cap x(0)=v\right). Further, event k≤ck\leq c can be partitioned and Eq. (14) can be written as

where second line is trivial since the events c=jc=j are disjoint. We can now use Bayes’ rule to derive the probability of uu being visited kk steps after vv and being selected in vv’s sampled context, as:

Now, let EvkuE_{vku} be the event that a walker visits vv and after kk steps, visits uu and selects it part of its context. This event happens with the probability indicated in Equation 18. Concretely,

Let Ev∗uE_{v*u} count the events {Evku:k∈[1,C]}\left\{E_{vku}:k\in[1,C]\right\}, then:

Suppose we run DeepWalk, starting mm random walks from each node vv, then the expected number of times that uu is present in the context of vv is given by:

Finally, we can write down the expectation over the square matrix D\mathbf{D}:

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: a→b→c  ⟹  a→ca\rightarrow b\rightarrow c\implies a\rightarrow c, 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.