Theory of Graph Neural Networks: Representation and Learning
Stefanie Jegelka
Introduction
Solving these tasks demands a sufficiently rich embedding of the graph or of each node that captures structural properties as well as the attribute information. While graph embeddings have been a widely studied topic, including spectral embeddings and graph kernels, recently, Graph Neural Networks (GNNs) have emerged as an empirically broadly successful model class that, as opposed to, e.g., spectral embeddings, allows to adapt the embedding to the task at hand, generalizes to other graphs of the same input type, and incorporates attributes. Due to space limits, this survey focuses on the popular message passing (spatial) GNNs, formally defined below, and a selection of their rich mathematical connections, with an excursion into higher-order GNNs.
When analyzing learning and risk, three main questions become important:
1. Representational power (Section 2). Which target functions can be approximated well by a GNN model class ? Answers to this question relate to graph isomorphism testing, approximation theory for neural networks, local algorithms and representing invariance/equivariance under permutations.
2. Generalization (Section 3). Even with sufficient approximation power, we can only estimate a function from : since and hence is not accessible, the common learning or training procedure is to minimize the empirical risk :
Generalization asks how well is performing according to the population risk, i.e., , as a function of and model properties. Good generalization may demand explicit (e.g., via penalties) or implicit regularization (e.g., via the optimization algorithm, typically variants of stochastic gradient descent). Hence, generalization analyses involve the complexity of the model class , the target function, the data and the optimization procedure.
3. Generalization under distribution shifts (Section 4). In practice, a learned model is often deployed on data from a distribution , e.g., graphs of different size, degree or attribute ranges, for instance, . In which cases can we expect successful extrapolation to ? This depends on the structure of the graphs and the task, formalizable via graph limits, local structures and algorithmic structures, e.g., dynamic programming.
Beyond these topics, GNNs have close connections to graph signal processing as learnable filters, geometric learning, and probabilistic inference. The first and third connection are covered in .
Lastly, a disclaimer: GNNs are a rapidly evolving research area. Hence, it is almost impossible to include all possible works, and any survey necessarily misses some. The author apologizes in advance for any references that were not covered.
The final node representation is the last iterate, possibly concatenated with a linear classifier. Throughout this paper, denotes the neighborhood of , and a multiset. Here, encodes the -hop neighborhood of node , i.e., the subgraph of all nodes reachable from within steps. The number of iterations is also termed the GNN depth, and one iteration may be viewed as a layer.
The sum may also be replaced by an average, degree-normalized sum or coordinate-wise min or max. In the most general form, the functions are implemented as multi-layer perceptrons (MLPs), neural networks that alternate linear transformations and coordinate-wise nonlinear activations such as the ReLU () or sigmoid function ():
The learnable parameters of the MLP are the weight matrices and bias vectors . The update in Equation (1.5) is typically a weighted combination with learnable weight matrices:
Finally, if a graph-level prediction is desired, all node representations can be aggregated by a permutation invariant readout function
Here, we assume the readout has the form (1.6) or is a simple sum or average. Typically, all parameters are learned jointly via stochastic gradient descent minimizing the empirical risk.
Throughout this article, denotes the number of nodes and the number of training data points.
Spectral GNNs. Besides message passing GNNs, other architectures have been devised. One important example are spectral graph neural networks , which learn a function of the graph Laplacian , i.e., , where and are the matrices of eigenvectors and eigenvalues (diagonal matrix), respectively, and is a polynomial. In the sequel, we will focus on message passing GNNs.
Representational power of GNNs
For functions on graphs, representational power has mainly been studied in terms of graph isomorphism: which graphs a GNN can distinguish. Via variations of the Stone-Weierstrass theorem, these results yield universal approximation results. Other works bound the ability of GNNs to compute specific polynomials of the adjacency matrix and to distinguish graphons . Observed limitations of MPNNs have inspired higher-order GNNs (Section 2.3). Moreover, if all node attributes are unique, then analogies to local algorithms yield algorithmic approximation results and lower bounds (Section 2.2).
A standard characterization of the discriminative power of GNNs is via the hierarchy of the Weisfeiler-Leman (WL) algorithm for graph isomorphism testing, also known as color refinement or vertex classification , which was inspired by the work of Weisfeiler and Leman . The WL algorithm does not entirely solve the graph ismomorphism problem, but its power has been widely studied.
A labeled graph is a graph endowed with a node coloring for some sufficiently large alphabet . Given a labeled graph , the 1-dimensional WL algorithm (1-WL) iteratively computes a node coloring . Starting with in iteration , in iteration it sets for all
where Hash is an injective map from the input pair to , i.e., it assigns a unique color to each neighborhood pattern. To compare two graphs , the algorithm compares the multisets and in each iteration. If the sets differ, then it determines that . Otherwise, it terminates when the number of colors in iteration and are the same, which occurs after at most iterations.
The computational analogy between the 1-WL algorithm and MPNNs is obvious. Since the WL algorithm uniquely colors each neighborhood, the coloring always refines the coloring from a GNN.
If for two graphs a message passing GNN outputs , then the 1-WL algorithm will determine that . For any , there exists an MPGNN such that . A sufficient condition is that the aggregate, update and readout operations are injective multiset functions.
GNNs that use the degree for normalization in the aggregation can be equivalent to the 1-WL agorithm too, but with one more iteration in the WL algorithm .
Theorem 1 demands the neighbor aggregation to be an injective multiset function on sets (). Theorem 2 shows how to universally approximate multiset functions.
The proof idea is to show that there exists an injective function of the form . The above result is an extension of a universal approximation result for set functions , and suggests a neural network model for sets where are approximated by MLPs. The Graph Isomorphism Network (GIN) implements this sum decomposition in the aggregation function to ensure the ability of injective operations.
1.2 Implications for graph distinction
Theorem 1 allows to directly transfer any known result for the 1-WL test to MPNNs. For instance, the 1-WL test succeeds to distinguish graphs sampled uniformly from all graphs on nodes with high probability, and failure probability going to zero as . 1-WL can also distinguish any non-isomorphic pair of trees . It fails for regular graphs, as all node colors will be the same. The graphs that the 1-WL algorithm can distinguish from any non-isomorphic graph can be recognized in quasi-linear time . See also for more detailed results on the expressive power of variants of the WL algorithm.
1.3 Computation trees and structural graph properties
To further illustrate the implications of GNNs’ discriminative power, we look at some specific examples. The maximum information contained in any embedding can be characterized by a computation tree , i.e., an “unrolling” of the message passing procedure. The 1-WL test essentially colors computation trees. The tree is constructed recursively: let for all . For , construct a root with label and, for any construct a child subtree . Figure 1 illustrates an example.
If for two nodes , we have , then .
Comparing computation trees directly implies that MPNNs cannot distinguish regular graphs. It also shows further limitations with practical impact, as indicated in Figure 1, in particular for learning combinatorial algorithms and for predicting properties of molecules, where functional groups are of key importance. We say a class of models decides a graph property if there exists an such that for any two that differ in the property, we obtain .
MPNNs cannot decide girth, circumference, diameter, radius, existence of a conjoint cycle, total number of cycles, and existence of a -clique . MPNNs cannot count induced (attributed) subgraphs for any connected pattern of 3 or more nodes, except star-shaped patterns .
Motivated by these limitations, generalizations of GNNs were proposed that provably increase their representational power. Two main directions are to (1) introduce node IDs (Section 2.2), and (2) use higher-order functions that act on tuples of nodes (Section 2.3).
2 Node IDs, local algorithms, combinatorial optimization and lower bounds
The major weaknesses of MPNNs arise from their inability to identify nodes as the origin of specific messages. Hence, MPNNs can be strengthened by making nodes more distinguishable. The gained representational power follows from connections with local algorithms, where the input graph defines both the computational problem and the network topology of a distributed system: there, each node is a local machine and generates a local output, and all nodes execute the same algorithm, without faults.
Approximation Algorithms. Sato et al. achieve a partial node distinction by transferring the idea of a port numbering from local algorithms. Edges incident to each node are numbered as outgoing ports. In each round, each node simultaneously sends a message to each port, but the messages can differ across ports:
Permutation invariance, though, is not immediate. This corresponds to the vector-vector consistent (VVC) model for local algorithms . The VVC analogy allows to transfer results on representing approximation algorithms. CPNGNN is a specific VVC GNN model.
There exists a CPNGNN that can compute a -approximation for the minimum dominating set problem, a CPNGNN that can compute a 2-approximation for the minimum vertex cover problem, but no CPNGNN can do better. No CPNGNN can compute a constant-factor approximation for the maximum matching problem.
Adding a weak vertex 2-coloring leads to further results. Despite the increased power compared to MPNN, CPNGNNs retain most limitations of Proposition 2 .
A more powerful alternative is to endow nodes with fully unique identifiers . E.g., augmenting the GIN model (the most expressive MPNN) with random node identifiers yields a model that can decide subgraphs that MPNN and CPNGNN cannot . This model can further achieve better approximation results for minimum dominating set (), where is the harmonic number) and maximum matching ().
Turing completeness. Analogies to local algorithms further imply that MPNNs with unique node IDs are Turing complete, i.e., they can compute any function that a Turing machine can compute, including graph isomorphism. In particular, the proof shows an equivalence to the Turing universal LOCAL model from distributed computing .
If and are Turing complete functions and the message passing GNN gets unique node IDs, then the classes GNN and LOCAL are equivalent. For any MPNN there exists a local algorithm of the same depth, such that , and vice versa.
Lower bounds. For bounded size, GNNs lose computational power. Via analogies to the CONGEST model , which bounds message sizes, one can transfer results on decision, optimization and estimation problems on graphs. These lead to lower bounds on the product of depth and width of the GNN if the nodes do not have access to a random generator. Here, as before, width of a GNN refers to the dimensionality of the embeddings .
If a problem cannot be solved in less than rounds in CONGEST using messages of at most bits, then it cannot be solved by an MPNN of width and depth , where .
Theorem 5 directly implies lower bounds for solving combinatorial problems, e.g., for cycle detection and computing diameter, and for minimum spanning tree, minimum cut and shortest path .
Random node IDs. While unique node IDs are powerful in theory, in many practical examples the input graphs do not have unique IDs. An alternative is to assign random node IDs . This can still yield GNNs that are essentially permutation invariant: while their outputs are random, the outputs for different graphs are still sufficiently separated . This leads to a probabilistic universal approximation result:
The proof builds on a result by that states that any logical sentence in FOC2 can be expressed by the addressed GNN. The logic considered here is a fragment of first-order (FO) predicate logic that allows to incorporate counting quantifiers of the form , i.e., there are at least elements satisfying , but is restricted to two variables. FOC2 is tightly linked with the 1-WL test: for any nodes in any graph, 1-WL colors and the same if and only if they are classified the same by all FOC2 classifiers .
Another successful idea is to augment node attribute vectors with attributes that contain structural or topological information .
3 Higher-order GNNs
Instead of adding unique node IDs, one may increase the expressive power of GNNs by encoding subsets of that are larger than the single nodes used in MPNNs. Three such directions are: (1) neural network versions of higher-dimensional WL algorithms, (2) (non)linear equivariant operations, and (3) recursion. Other strategies that could not be covered here use, e.g., simplicial and cell complexes .
Extending analogies between MPNNs and the 1-WL algorithm , the first class of higher-order GNNs imitates versions of the -dimensional WL algorithm. The -WL algorithms are defined on -tuples of nodes, and different versions differ in their aggregation and definition of neighborhood. In iteration , the -WL algorithm labels each -tuple by a unique ID for its isomorphism type. Then it aggregates over neighborhoods for :
For two graphs the -WL algorithm then decides “not isomorphic” if for some , and returns “maybe isomorphic” otherwise. Like the 1-WL test, the -WL test decides “not isomorphic” only if . The Folklore -WL algorithm (-FWL) differs in its update rule, which “swaps” the order of the aggregation steps :
The 1-WL and 2-WL tests are equivalent, and for , the -WL test can distinguish strictly more graphs than the -WL test . The -FWL algorithm is as powerful as the -WL algorithm for .
Set-WL GNN. Since computations on -tuples are expensive, Morris et al. consider a GNN that corresponds to a set version of a -WL algorithm. For any set with , let . The set-based WL (-SWL) algorithm then updates as
its GNN analogue uses the aggregation and update (cf. Eqns. (1.6),(1.8))
where is a coordinatewise nonlinearity (e.g., sigmoid or ReLU). This family of GNNs is equivalent in power to the -SWL test (Theorem 8). For computational efficiency, a local version restricts the neighborhood of to sets such that the nodes in the symmetric difference are connected in the graph. This local version is weaker .
Folklore WL GNN. In analogy to the -FWL algorithm, Maron et al. define -FGNNs with aggregations
The family of -FGNNs is a class of nonlinear equivariant networks, and is equivalent in power to the -FWL test and the -WL test (Theorem 8).
3.2 Linear equivariant layers.
The associated GNN model uses one parameter (coefficient) for each basis tensor. Importantly, the number of parameters is independent of the number of nodes. The proof for identifying the basis tensors sets up a fixed point equation with Kronecker products of any permutation matrix that any equivariant tensor must satisfy. The solutions to these equations are defined by equivalence classes of multi-indices in . Each equivalence class is represented by a partition of , e.g., includes all multi-indices where and . The basis tensors are then such that if and only if .
Linear equivariant GNNs of order (-LEGNNs) parameterized with the full basis are as discriminative as the -WL algorithm (Theorem 8). To achieve this discriminative power, each entry in the input tensor encodes an initial coloring of the isomorphism type of the subgraph indexed by the -tuple .
3.3 Summary of Representational Power via WL
The following theorem summarizes equivalence results between the GNNs discussed so far and variants of the WL test. Following , we here use equivalence relations, as they suffice for universal approximation in Section 2.4. For a set of functions defined on , define an equivalence relation via the joint discriminative power of all functions , i.e., for any two graphs :
The above GNN families have the following equivalences:
Analogous results hold for equivariant models (for node representations), with the exception of equality (2.15), which becomes an inclusion: \rho(\text{k-LGNN}_{E})\subseteq\rho(\text{k-WL}_{E}) .
3.4 Relational Pooling.
where is with permuted rows, and is the tensor combining adjacency matrix and node attributes. Here, is any permutation-sensitive function, and may be modeled via various nonlinear function approximators, e.g. neural networks such as fully connected networks (MLPs), recurrent neural networks or a combination of a convolutional network applied to and an MLP applied to . In particular, this model allows to implement graph isomorphism testing via node IDs (cf. Section 2.2) if is a universal approximator . For instance, node IDs may be permuted over nodes and concatenated with the node attributes:
The drawback of the Relational Pooling (2.17) is its computational intractability. Various approximations have been considered, e.g., defining canonical orders, stochastic approximations, and applying to all possible -subsets of . In the latter case, increasing strictly increases the expressive power. Local Relational Pooling is a variant that applies relational pooling to the -hop subgraphs centered at each node, and then aggregates the results. This operation provably allows to identify and count subgraphs of size up to .
3.5 Recursion
A general strategy for encoding a graph is to encode a collection of subgraphs and then aggregate these encodings. The question of what graphs this process allows to distinguish (), depends on the collection of subgraphs used, the subgraph encoding function and the aggregation function. As a special case, this process includes the reconstruction hypothesis , i.e., the question whether any graph can be reconstructed from the collection of its subgraphs , for all in . One challenge with the reconstruction hypothesis is that no alignment of the is available. Indeed, node correspondences across subgraphs provide important information .
Indeed, the expressive power of a model based on subgraph encodings depends on the set of subgraphs, the type of subgraph encodings and the aggregation. Tahmasebi et al. use recursion as a powerful tool: instead of iterative message passing or layering, a recursive application of the above subgraph embedding step, even with a simple set aggregation like (1.6), can enable a GNN that can count any bounded-size subgraphs, as opposed to MPNNs (Prop. 2).
Let be the -hop neighborhood of in . Recursive neighborhood pooling (RNP) encodes intersections of such neighborhoods of different radii. Given an input graph with node attributes and a sequence of radii, RNP-GNN recursively encodes the node-deleted -neighborhoods of all nodes after marking the deletion in augmented representations , . It then combines the results, and returns node representations of all nodes. Concretely, for each , it computes and
If the sequence of radii is empty (base case), then the algorithm returns the input attributes . In contrast to iterative message passing, the encoded subgraphs here correspond to intersections of local neighborhoods. Together with the node deletions and markings that retain node correspondences, this maintains more structural information. Formally, if the sequence of radii dominates a covering sequence for a subgraph of interest, then, with appropriate parameters, RNP can count the induced and non-induced subgraphs of isomorphic to . The computational cost is for recursion depth , and better for very sparse graphs, in line with computational lower bounds.
4 Universal approximation
Distinguishing given graphs is closely tied to approximating continuous functions on graphs. In early work, Scarselli et al. take a fixed point view and show a universal approximation result for infinite-depth MPNNs whose layers are contraction operators, for functions on equivalence classes defined by computation trees. Dehmamy et al. analyze the ability of GNNs to compute polynomials of the adjacency matrix.
Later works derive universal approximation results for graph and permutation-equivariant functions from graph discrimination results via extensions of the Stone-Weierstrass theorem . Maron et al. argue that -invariant networks (for a permutation group ) can universally approximate -invariant polynomials, which in turn can universally approximate any invariant function . Keriven and Peyré do not fix the size of the graph and show that shallow equivariant networks can, with a single set of parameters, well approximate a function on graphs of varying size. Both constructions involve very large tensors.
Generalization
Beyond approximation power, a second important question in machine learning is generalization. Generalization asks how well the estimated function is performing according to the population risk, i.e., , as a function of the number of data points and model properties. Good generalization may demand explicit (e.g., via a penalty term) or implicit regularization (e.g., via the optimization algorithm). Hence, generalization analyses involve aspects of the complexity of the model class , the target function we aim to learn, the data and the optimization procedure. This is particularly challenging for neural networks, due to the nested functional form and the non-convexity of the empirical risk.
A classic learning theoretic perspective bounds the generalization gap via the complexity of the model class (Section 3.1). These approaches do not take into account possible implicit regularization via the optimization procedure. One possibility to do so is via the Neural Tangent Kernel approximation (Section 3.2). Finally, for more complex, structured target functions, e.g., algorithms or physics simulations, one may want to also consider the structure of the target task. One such option is Algorithmic Alignment (Section 3.3). Another strategy for obtaining generalization bounds is via algorithmic stability, the condition that, if one data point is replaced, the outcome of the learning algorithm does not change much. This strategy led to some early bounds for spectral GNNs .
Vapnik-Chervonenkis dimension. The first GNN generalization bound was based on bounding the Vapnik-Chervonenkis (VC) dimension of the GNN function class . The VC dimension of expresses the maximum cardinality of a set of data points such that for any binary labeling of the data, some GNN in can perfectly fit, i.e., shatter, the set. The VC dimension directly leads to a bound on the generalization gap. Here, we only state the results for sigmoid activation functions.
The VC dimension of GNNs with parameters, hidden neurons (in the MLP) and input graphs of size is .
Strictly speaking, Theorem 10 is for node classification with one hidden layer in the aggregation function MLPs. The VC dimension directly yields a bound on the generalization gap: for a class with VC dimension , with probability , it holds that
Interestingly, in these bounds, GNNs are a generalization of recurrent neural networks . The VC dimension bounds for GNNs are the same as for recurrent neural networks ; the bounds for fully connected MLPs are missing the factor .
Rademacher Complexity. Bounds that are in many cases tighter can be obtained via Rademacher complexity. The empirical Rademacher complexity of a function class measures how well it can fit “noise” in the form of uniform random variables in :
Let be the product of the Lipschitz constants of and ; the number of GNN iterations; the dimension of the embeddings , and the maximum branching factor in the computation tree. Then the generalization gap of the GNN can be bounded as: for , for and for .
2 Generalization bounds via the Neural Tangent Kernel
In contrast to the results in Section 3.1, the complexity measure of the target function is data-dependent. If the target function to be learned follows a simple GNN structure with a polynomial, then this bound can be polynomial:
Let . If the labels , , satisfy
3 Generalization via Algorithmic Alignment
The Graph NTK analysis shows a polynomial sample complexity if the function to be learned is close to the computational structure of the GNN, in a simple way. While this applies to mainly simpler learning tasks, the idea of an “alignment” of computational structure carries further. Recently, there has been growing interest in learning scientific tasks, e.g., given a set of particles or planets along with their location, mass and velocity, predict the next state of the system , and in “algorithmic reasoning”, e.g., learning to solve combinatorial optimization problems in particular over graphs . In such cases, the target function corresponds to an algorithm, e.g., a dynamic program.
While many neural network architectures have the power to represent such tasks, empirically, they do not learn them equally well from data. In particular, GNNs perform well here, i.e., their architecture encodes suitable inductive biases . As a concrete example, consider the Shortest Path problem. The computational structure of MPNNs matches that of the Bellman-Ford (BF) algorithm very well: both “algorithms” iterate, and in each iteration , update the state as a function of the neighboring nodes and edge weights :
Hence, the GNN can simulate the BF algorithm if it uses sufficiently many iterations, and if the aggregation function approximates the BF state update (relaxation step). Intuitively, this update is a much simpler function to learn than the full algorithm as a black box, i.e., the GNN encodes much of the algorithmic structure, sparsity and invariances in the architecture. More generally, MPNNs match the structure of many dynamic programs in an analogous way , as long as the updates are permutation invariant or sufficient node identification is provided as input, in light of the results in Section 2. Dudzik and Veličković refine and generalize the relations between GNNs and dynamic programming by using category theory.
The NTK results formalize simplicity by a small function norm in the RKHS associated with the Graph NTK; this can become complicated with more complex tasks and multiple layers. To quantify structural match, Xu et al. define algorithmic alignment by viewing a neural network as a structured arrangement of learnable modules – in a GNN, the (MLPs in the) aggregation functions – and define complexity via sample complexity of those modules in a PAC-learning framework. Sample complexity in PAC learning is defined as follows: We are given a data sample drawn i.i.d. from a distribution that satisfies for an underlying target function . Let be the function output by a learning algorithm . For a fixed error and failure probability , the function is -PAC learnable with if
The sample complexity is the smallest so that is -learnable with .
Let be a target function and a neural network with modules . The module functions generate for if, by replacing with , the network simulates . Then -algorithmically aligns with if (1) generate and (2) there are learning algorithms for learning with , with sample complexity .
Algorithmic alignment resembles Kolmogorov complexity . Thus, it can be hard to obtain the optimal alignment between a neural network and an algorithm. But, any algorithmic alignment yields a bound, and any with acceptable sample complexity may suffice. The complexity of the MLP modules in GNNs may be measured with a variety of techniques. One option is the NTK framework. The module-based bounds then resemble the polynomial bound in Theorem 13, since both are extensions of . However, here, the bounds are applied at a module level, and not for the entire GNN as a unit. Theorem 14 translates these bounds, in a simplified setting, into sample complexity bounds for the full network.
Fix and . Suppose , where , and for some . Suppose are network ’s MLP modules in sequential order of processing. Suppose and -algorithmically align via functions for a constant . Under the following assumptions, is -learnable by . a) Sequential learning. We train ’s sequentially: has input samples , with obtained from . For , the input for are the outputs of the previous modules, but labels are generated by the correct functions on . b) Algorithm stability. Let be the learning algorithm for the ’s. Suppose , and . For any , , for some . c) Lipschitzness. The learned functions satisfy , for some .
The big notation here hides factors including the Lipschitz constants, number of modules and graph size. When measuring module complexity via the NTK, Theorem 14 indeed yields a gap in upper bounds between fully connected networks and GNNs in simple cases , supporting empirical results. While some works use sequential training in experiments , empirically, better alignment improves learning and generalization even with more common “end-to-end” training, i.e., optimizing all parameters simultaneously .
At a general level, these alignment results indicate how incorporating expert knowledge, e.g. in terms of algorithmic techniques or physics, into the design of the learning method can improve sample efficiency.
Extrapolation
Structural similarity of graphs. One possibility to guarantee successful extrapolation to larger graphs is to assume sufficient structural similarity between the graphs in and , in particular, structural properties that matter for the GNN family under consideration. For spectral GNNs, this assumption has been formalized as the graphs arising from the same underlying topological space, manifold or graphon. Under such conditions, spectral GNNs – with conditions on the employed filters – can generalize to larger graphs . The underlying structure also ensures similar local structure of the graphs.
For spatial message passing GNNs, whose representations rely on computation trees as local structures (Section 2.1), an agreement in the distributions of the computation trees in the graphs sampled from and is necessary . This is violated, for instance, if the degree distribution is a function of the graph size, as is the case for random graphs under the Erdős-Rényi or Preferential Attachment models. The computation tree of depth rooted at a node corresponds to the color assigned by the 1-WL algorithm.
Let and be finitely supported distributions of graphs. Let be the distribution of colors over and similarly for . Assume that any graph in contains a node with a color in . Then, for any graph regression task solvable by a GNN with depth there exists a GNN with depth at most that perfectly solves the task on and predicts an answer with arbitrarily large error on all graphs from .
The proof exploits the fact that GNN predictions on nodes only depend on the associated computation tree and that a sufficiently flexible GNN can assign arbitrary target labels to any computation tree . I.e., the available information allows for multiple local minima of the empirical risk. A similar result can be shown for node prediction tasks. (A “sufficiently large” GNN here means depth at least layers and width , where the max degree refers to any graph in the support, is the finite number of possible input node attributes and the set of colors encountered in graphs in the support.)
Conditions on the GNN. If sufficient structural similarity of the input graphs cannot be guaranteed, then further restrictions on the GNN can enable extrapolation to different graph sizes, structures and ranges of input node attributes. If there are no training observations in a certain range of attributes or set of local structures, then the predictions of the learned model depend on the inductive biases induced by the model architecture, loss function and training algorithm. Which prediction function, out of multiple fitting functions, a model will choose, depends on these biases.
Xu et al. analyze such biases to obtain conditions on the GNN for extrapolation. Taking the perspective of algorithmic alignment (Section 3.3), they first analyze how individual module functions, i.e., the MLPs in the aggregation function of a GNN, extrapolate, and then transfer this to the entire GNN. The aggregation functions enter the extrapolation regime, e.g., if the node attributes, node degrees or computation trees are different under compared to , as they determine the inputs to the aggregations. The following theorem states that, sufficiently far away from , MLPs implement directionally linear functions.
The linear function and the constant terms in the convergence rate depend on the training data and the direction . The proof of Theorem 16 relies on the fact that a neural network in the NTK regime learns a minimum-norm interpolation function . Although Theorem 16 uses a simplified setting of a wide 2-layer network, similar results hold empirically for more general MLPs .
To appreciate the implications of this result in the context of GNNs, consider the example of Shortest Path in Equation (3.4). For the aggregation function to mimic the Bellman-Ford algorithm, the MLP must approximate a nonlinear function. But, in the extrapolation regime, it implements a linear function and therefore is expected to not approximate Bellman Ford well any more. Indeed, empirical works that successfully extrapolate GNNs for Shortest Path use a different aggregation function of the form
Here, the nonlinear parts do not need to be learned, allowing to extrapolate with a linear learned MLP. More generally, the directionally linear extrapolation suggests that (1) the architecture or (2) the input encoding should be set up such that the target function can be approximated when MLPs learn linear functions (linear algorithmic alignment). An example for (2) may be found in forecasting physical systems, e.g., predicting the evolution of objects in a gravitational system, and the node (object) attributes are mass, location and velocity at time . The position of an object at time is a nonlinear function of the attributes of the other objects. When encoding the nonlinear function as transformed edge attributes, the function to be learned becomes linear. Many empirical works that successfully extrapolate implement the idea of linear algorithmic alignment .
Finally, the geometry of the training data also plays an important role. show empirical results and initial theoretical results for learning max-degree, that, even with linear algorithmic alignment, sufficient diversity in the training data is needed to identify the correct linear functions. These data conditions are weaker than those implied by Theorem 15, due to the linear algorithmic alignment assumption.
For the case when the target test distribution is known, Yehudai et al. propose approaches for combining elements of and to enhance the range of the data seen by the GNN.
Conclusion
This survey summarized three main topics in theoretically understanding GNNs: representation and approximation, generalization, and extrapolation. As GNNs are an active research area, many results could not be included. E.g., we focused on MPNNs and main ideas for higher-order GNNs, but neglected spectral GNNs, which closely relate to ideas in graph signal processing. Other emergent topics include adversarial robustness, optimization behavior of the empirical risk and its improvements, and computational scalability and approximations. Overall, GNNs have a rich set of mathematical connections, a selection of which was covered here.
Many questions remain. Regarding approximation capabilities, the limitations of MPNNs have motivated powerful higher-order GNNs. However, these are still computationally expensive. What efficiency is theoretically possible? Moreover, most applications may not require full graph isomorphism power, or -WL power for large . What other measures of representational power make sense? Do they allow better and sharper complexity results? Initial works consider subgraph counting as a benchmark task .
The generalization results so far need to use simplifications in the analysis, similar to most theoretical analyses of deep learning. To what extent can they be relaxed? Do more specific tasks or graph classes allow sharper results? Which modifications of GNNs would allow them to generalize better, and how do higher-order GNNs generalize? Similar questions pertain to extrapolation and reliability under distribution shifts, a topic that has been studied even less than GNN generalization.
In general, revealing further mathematical connections may enable the design of richer models and enable a more thorough understanding of GNNs’ learning abilities and limitations, and eventual improvements.
The author would like to thank Keyulu Xu, Derek Lim, Behrooz Tahmasebi, Vikas Garg, Tommi Jaakkola, Andreas Loukas, Jingling Li, Mozhi Zhang, Simon Du, Ken-ichi Kawarabayashi, Weihua Hu, Jure Leskovec, Joan Bruna and Yusu Wang for discussions on the theory of GNNs, collaborations and pointers.
This work was partially supported by NSF CAREER award 1553284, NSF SCALE MoDL award 2134108, and NSF CCF-2112665 (TILOS AI Research Institute).