Theme
Graph Neural Networks
One-line definition: a graph neural network (GNN) is a neural network that runs directly on graph-structured data — it lets each node learn by "collecting neighbor information and updating its own representation," thereby encoding nodes, edges, and even entire graphs into low-dimensional vectors. If CNN is the neural network designed for "pixel grids" and RNN/Transformer for "sequences," then GNN is designed for relationships: every operation it performs is built on "who connects to whom."
The real world is full of graphs: social networks are "person–person" graphs, molecules are "atom–chemical bond" graphs, knowledge graphs are "entity–relation" graphs, and so are power grids, transportation networks, and protein interaction networks. This creates a huge gap: deep learning conquers images and text, but can't directly apply convolutional kernels or attention matrices to a "graph with no fixed grid, variable node count, and no fixed order." GNN fills precisely this gap — it brings deep learning into the era of relational data. Its mathematical foundation, parameter update rules, and supervised framework are exactly the same as what you learned in What is Machine Learning; only the data form goes from "matrix" to "graph."
1. Problem Definitions for Graph Data
1. What is a graph
A graph (Graph) consists of two parts:
Graph G = (V, E)
V: set of vertices / nodes, e.g., users, atoms, entities
E: set of edges, e.g., follow relationships, chemical bonds, relations in tripletsGraphs can be further classified by edge attributes:
| Type | Does the edge have direction? | Does the edge have weight? | Examples |
|---|---|---|---|
| Undirected unweighted | No | No | Friendships, protein interactions |
| Directed unweighted | Yes | No | Weibo follows, web links, citation relationships |
| Undirected weighted | No | Yes | Transportation networks (edge weight = distance), similarity networks |
| Directed weighted | Yes | Yes | Money transfers (direction + amount) |
| Heterogeneous | Nodes/edges have types | — | Knowledge graphs (entity types + relation types) |
Compared to images and text, graphs have three "anti-deep-learning-intuition" properties:
- No fixed grid: you can't put a node into "pixel (i,j)"; the graph structure itself is part of the data;
- Variable size: each graph has different numbers of nodes and edges;
- Permutation invariance: renumbering nodes doesn't change the graph — the model must be robust to such reordering.
2. Three classes of tasks
By "which level of the graph the prediction targets," supervised tasks on graphs fall into three categories:
| Task level | Prediction target | Typical question | Examples |
|---|---|---|---|
| Node-level | A single node | Node classification, node regression | Determine if a user in a social network is a bot; predict whether an atom affects molecular toxicity |
| Edge-level | A pair of nodes | Link prediction, edge classification | Predict whether two people will become friends; predict whether two drugs interact; knowledge graph completion |
| Graph-level | An entire graph | Graph classification, graph regression | Predict whether a molecule can cross the blood-brain barrier; predict whether a code function contains a vulnerability |
Node classification is the classic scenario for "semi-supervised": some nodes on the graph have labels, the rest don't, and the model uses graph structure to "spread" labels — precisely GNN's most capable ability. Link prediction is "treating edges as labels": remove a portion of edges from the graph as training labels, predict the remaining ones. Understanding these three tasks is essential for choosing the right model output layer: node classification uses softmax per node; graph classification reads out one vector from the entire graph and feeds it to a classification head.
2. Graph Representation: Adjacency Matrix and Features
The first step in getting deep learning to handle graphs is turning them into tensors. The standard representation is the adjacency matrix:
A B C D E
A [0 1 1 0 0]
B [1 0 0 1 0] A_ij = 1 ⇔ there's an edge between node i and j
C [1 0 0 1 0]
D [0 1 1 0 1] — undirected graph = symmetric matrix; weighted graph uses weights instead of 0/1
E [0 0 0 1 0]Two important objects paired with the adjacency matrix:
- Degree matrix D: a diagonal matrix,
D_ii = Σ_j A_ij, i.e., the number of neighbors of node i (its degree). Degree is a first-order measure of a node's "importance" in the network; hub nodes have degrees far above average. - Feature matrix X: shape n×d, row i is node i's d-dimensional features. The source of features varies by domain: social networks can be user profile vectors, molecules can be atom type one-hot (C, N, O...), academic networks can be paper bag-of-words vectors.
With A and X, a naive idea is to feed X directly to an MLP — but this wastes the graph structure: MLP only looks at "each node's own features," nodes are completely independent, and neighbor info is unused. Another naive idea is to concatenate A and X and feed to a fully-connected layer — but this has two fatal problems:
- Permutation sensitivity: renumbering nodes rearranges the matrix rows/columns; the same graph yields a completely different vector;
- Parameter can't be shared: different graphs have different node counts, so fixed-size weights can't be used.
All GNN design is about finding operations that "both use structure and are invariant to node reordering" under these two constraints. The good news is: "aggregate neighbor info" natively satisfies permutation invariance — summing or averaging neighbor features is independent of the traversal order of neighbors. GNN starts precisely from here.
3. Core Idea: Message Passing
1. Intuition
Imagine you want to learn a vector representation for each node on a graph, making it "know where it sits and who it connects to." A simple but powerful idea is to iterate:
Each round, every node sends its representation to neighbors and receives messages from all neighbors, aggregates them into a new message, then combines with its own representation to update itself.
This is like message diffusion in a village: in the first round, you only know what your neighbors say; in the second round, you know what "your neighbors' neighbors" say. The more rounds, the larger "view" each node has. This mechanism is called message passing, and it's the common skeleton of almost all GNNs:
Layer 0: ① ② ③ ④ Each node only has its own features h⁰
\ /│\ │/
Layer 1: ①′ ③′ ①′ = aggregate(②,③) combined with own features
/│\ /│\
Layer 2: … ②″ … ②″ has seen all info within two hops2. Formal definition
Let h_v^(l) be node v's representation at layer l, N(v) be v's neighbor set. One layer of message passing can be written as a universal template:
h_v^(l+1) = UPDATE( h_v^(l) , AGGREGATE( { h_u^(l) : u ∈ N(v) } ) )
↑ update self ↑ aggregate neighbors' messages
Layer l+1 function aggregation function (sum/mean/max...)The two layers have different responsibilities:
- AGGREGATE: merge neighbors' vectors into one vector. It must be order-insensitive to neighbors (permutation invariant) — sum, mean, and max all satisfy this, while "concatenate in order" doesn't.
- UPDATE: fuse "old self representation" and "aggregated neighbor info" into a new representation, usually by concatenating and passing through a fully-connected layer with a nonlinear activation.
3. Stacking layers = receptive field expansion
The most critical sentence for understanding GNN depth:
After a k-layer GNN, node v's representation encodes the entire subgraph within radius k hops centered at v.
Analogy with CNN receptive field:
| CNN | GNN | |
|---|---|---|
| Basic operation | Convolutional kernel scans over pixel neighborhoods | Aggregation function scans over graph neighborhoods |
| Receptive field | Kernel size × layers | Hop count (1 layer = 1 hop of neighbors) |
| Data form | Fixed grid | Arbitrary topology |
| Weight sharing | Convolutional kernel shared across the whole graph | Transformation matrix W shared across all nodes |
The first layer can only see direct neighbors; the second layer can see two-hop (through intermediate nodes); the more layers, the more "big picture info" each node's representation contains. All representative models — GCN, GraphSAGE, GAT, GIN, MPNN — are just different instantiations of the AGGREGATE and UPDATE in the template above. This unified perspective comes from the MPNN (Message Passing Neural Network) framework proposed by Gilmer et al. in 2017, which decomposes aggregation into "message function + aggregation function + update function," proving that a broad class of graph models is just different configurations of the same framework.
4. Minimal implementation
Stripping away all wrappers, message passing itself is only a few dozen lines of code (PyTorch style):
python
import torch
def message_passing_round(h, adj, W, sigma):
"""h: n×d node features; adj: n×n adjacency matrix; W: transformation matrix"""
msgs = adj @ h # ① Each node receives the sum of neighbor features (message)
out = sigma(msgs @ W) # ② Transform + nonlinearity (update)
return out # One GNN layerThe adj @ h line is "aggregate neighbors" — matrix multiplication sums each node's neighbor features. It's efficient, parallelizable, and natively satisfies permutation invariance. The GCN, GraphSAGE, and GAT sections that follow all center on "how to do these two steps better."
4. GCN: Bringing Spectral Methods to Practice
GCN (Graph Convolutional Network) was proposed by Kipf & Welling in 2017 and is the most widely disseminated and classic GNN. Understanding it requires a derivation chain from "spectral graph theory" to "first-order approximation."
1. The idea of spectral methods: transforming the graph to the frequency domain
Why can image convolution locally smooth and denoise? Because of the convolution theorem: convolution in time domain = element-wise multiplication in frequency domain. The graph's spectral methods bring the same idea to graphs: for the graph Laplacian
L = D - A (Laplacian matrix, the "discrete second derivative" in graph theory)do spectral decomposition L = UΛUᵀ, where the columns of U are the graph's "frequency domain basis." Thus graph signal convolution can be written: first project signal x to the frequency domain (x̂ = Uᵀx), scale pointwise in the frequency domain (multiply by filter gθ), then project back to the spatial domain. Defferrard et al. (2016) approximated this filter with Chebyshev polynomials, avoiding expensive eigendecomposition — a key step making "spectral methods practical."
2. GCN: the ingenuity of first-order approximation
GCN's core insight is to limit the filter to first-order, paired with two tricks, yielding an exceptionally concise propagation rule:
H^(l+1) = σ( Â · H^(l) · W^(l) )
where:
à = A + I — add self-loops to each node (so updates "remember self")
D̃ = diag(Σⱼ Ãᵢⱼ) — degree matrix after adding self-loops
 = D̃^(-1/2) · à · D̃^(-1/2) — symmetrically normalized adjacency matrixBreaking down the meaning of each step:
- Self-loop
A + I: otherwise the proportion of "self info" in the aggregation result varies with degree; adding self-loops ensures each node always retains its own representation in updates; - Symmetric normalization
D̃^(-1/2)(·)D̃^(-1/2): each node's contribution is scaled by degree — high-degree hub nodes' messages are diluted, low-degree nodes amplified, preventing the aggregation result from being dominated by degree and also preventing numerical instability; H·W: each node transforms its features through the same weight matrix W — this is precisely "weight sharing," consistent with CNN convolutional kernels being shared across the whole graph.
Thus each GCN layer basically does "weighted average of neighbor features (weights determined by structure), then pass through a linear transform and nonlinear activation." Taking node i's update as an example, equivalently writing in component form:
h_i^(l+1) = σ( Σ_{j ∈ N(i) ∪ {i}} (1/√(d̃ᵢ·d̃ⱼ)) · W·h_j^(l) )python
import torch
import torch.nn as nn
def normalize(A):
"""Construct symmetrically normalized adjacency matrix Â"""
A_tilde = A + torch.eye(A.shape[0])
d = A_tilde.sum(dim=1)
D_inv_sqrt = torch.diag(d ** -0.5)
return D_inv_sqrt @ A_tilde @ D_inv_sqrt
class GCNLayer(nn.Module):
def __init__(self, in_dim, out_dim):
super().__init__()
self.W = nn.Linear(in_dim, out_dim)
def forward(self, x, A_hat):
return A_hat @ self.W(x) # Aggregate (structure-weighted) + feature transform
class GCN(nn.Module):
def __init__(self, nfeat, nhid, nclass):
super().__init__()
self.layer1 = GCNLayer(nfeat, nhid)
self.layer2 = GCNLayer(nhid, nclass)
def forward(self, x, A_hat):
h = torch.relu(self.layer1(x, A_hat))
return torch.log_softmax(self.layer2(h, A_hat), dim=1)3. Semi-supervised node classification
GCN's classic setting is semi-supervised node classification: on citation networks like Cora, CiteSeer, PubMed, each paper is a node, citation relationships are edges, and paper topics are labels, but only a few nodes have labels (only 20 per class in Cora). GCN's approach is to train on the entire graph: even nodes without labels participate in forward propagation (their neighbors have labels, and backpropagation sends gradients to these nodes), effectively using structural information as a "label diffusion" channel. The paper reports 81.5% accuracy on Cora, significantly exceeding earlier feature-based methods. This setting remains one of the standard baselines for evaluating GNNs to this day.
An intuition
Why can "structure" itself provide supervision? In citation networks, paper topics cluster highly locally — papers on the same topic cite each other. So "what topic are my neighbors on?" is a very strong signal for "what am I." What GNN learns is encoding this homophily pattern into representations. This also foreshadows a later pitfall: when the graph is heterophilic (e.g., fraud networks, where neighbors tend to be different classes), naive GNNs fail.
5. GraphSAGE: Sampling and Induction
GCN has two practical pain points: ① the entire training process needs the full-graph adjacency matrix; the larger the graph, the more memory explodes; ② it's transductive — only nodes seen during training can be predicted; new nodes (e.g., newly registered users) can't be inferred immediately. GraphSAGE (Hamilton et al., NeurIPS 2017) solves both problems simultaneously.
1. Inductive learning and neighbor sampling
GraphSAGE's core change: don't aggregate all neighbors, but randomly sample a fixed number (e.g., 25) of neighbors each time. The benefits:
- Memory is controllable: each step only needs a small subgraph, so even very large graphs can be trained;
- Inductive ability: the model learns "how to aggregate neighbors" as a function, not "the representation of a specific node" — as long as a new node's features are computed and edges connected, it can be inferred immediately. This lets GraphSAGE process graphs completely unseen during training (this is inductive, as opposed to transductive);
- Variance as regularization: sampling itself introduces randomness, acting as an implicit regularization.
2. Three aggregators
GraphSAGE's update rule adds "concat + normalize" to the universal template:
h_N(v)^(k) = AGGREGATE_k( { h_u^(k-1) : u ∈ N_sample(v) } )
h_v^(k) = σ( W^(k) · CONCAT( h_v^(k-1), h_N(v)^(k) ) )
h_v^(k) = h_v^(k) / ‖h_v^(k)‖₂ ← L2 normalization, stabilizes training"Self" and "neighbors" are encoded separately then concatenated, letting the model explicitly distinguish "who am I" from "who surrounds me." The paper compared three aggregation functions:
| Aggregator | Approach | Characteristics |
|---|---|---|
| Mean | Average neighbor vectors | Simple and fast; equivalent to a simplified GCN |
| LSTM | Shuffle neighbors randomly then pass through LSTM | Stronger expressiveness, but sacrifices permutation invariance (compensated by random shuffling); slower |
| Pool (max-pooling) | Each neighbor passes through a shared MLP, then take max per dimension | Balance between expressiveness and efficiency; commonly used in practice |
3. Training objectives
GraphSAGE supports both supervised and unsupervised training. The unsupervised loss uses graph adjacency as a self-supervised signal: make representations of connected nodes closer and representations of randomly sampled negative nodes farther apart (essentially a skip-gram-style objective). Representations learned this way can be reused for various downstream tasks, representing an early practice of "graph pre-training."
python
def sage_update(h_v, h_neighbors, W, sigma):
h_agg = h_neighbors.mean(dim=0) # Mean aggregation
h_new = sigma(W @ torch.cat([h_v, h_agg])) # Concat + transform
return h_new / h_new.norm() # L2 normalizationThe most famous engineering variant in the GraphSAGE family is PinSage — Pinterest uses it for recommendation on large-scale product-user graphs (see Application Scenarios), achieving industrial deployment at the billion-edge scale.
6. GAT: Attention Mechanism Takes Over Weights
GCN's aggregation weights are entirely determined by graph structure: 1/√(d̃ᵢd̃ⱼ) depends only on the degrees of two nodes. But in real graphs, neighbors aren't equally important — among all a user's friends, only one may have interests highly aligned with yours. GAT (Graph Attention Network, Veličković et al., ICLR 2018) lets weights be determined by content: just as Transformer learns "which words to pay attention to," GAT learns "which neighbors to pay attention to."
1. Computing attention coefficients
For each neighbor j of node i, GAT first computes a "relevance score," then normalizes across the neighbor set with softmax:
e_ij = LeakyReLU( aᵀ · [ W·h_i ‖ W·h_j ] ) — concatenate then score
α_ij = exp(e_ij) / Σ_{k∈N(i)} exp(e_ik) — normalize within neighbors
h_i' = σ( Σ_{j∈N(i)} α_ij · W·h_j ) — weighted aggregation by attentionwhere a is a learnable attention vector, ‖ denotes concatenation. Note that normalization happens within each node's neighbor set (Σ_k∈N(i)), not the whole graph — this ensures the same "each node only cares about its own neighborhood" locality as GCN.
2. Multi-head attention
GAT directly borrows Transformer's multi-head trick: use K independent attention heads to compute in parallel, concatenate results in the middle layers, and average in the output layer:
h_i' = ‖_{k=1..K} σ( Σ_{j∈N(i)} α_ij^(k) · W^(k) · h_j ) (middle layers: concatenate)
h_i' = σ( (1/K) Σ_{k=1..K} Σ_{j∈N(i)} α_ij^(k) · W^(k) · h_j ) (output layer: average)Multiple heads are like "looking at neighbors from multiple relational perspectives simultaneously," improving stability and expressiveness.
python
def gat_update(h_i, h_neighbors, W, a, sigma):
Wh_i = W @ h_i
Wh_j = W @ h_neighbors
e = (a @ torch.cat([Wh_i, Wh_j], dim=1)).squeeze() # Score
alpha = torch.softmax(e, dim=0) # Normalize within neighbors
return sigma((alpha * Wh_j).sum(dim=0)) # Weighted sum3. GCN vs GAT: a clear table
| Dimension | GCN | GAT |
|---|---|---|
| Weight source | Graph structure (degree) | Content (attention) |
| Weights learnable? | No, precomputed | Yes, optimized during training |
| Depends on adjacency matrix? | Yes (needs Â) | No, only needs edge list |
| Differentiation between neighbors | Only by degree | Fine-grained by content |
| Computational cost | Lower | Slightly higher with multi-head |
| Suitable scenario | Baseline, homophilic graphs | Graphs where neighbor importance varies significantly |
An empirical pattern: GAT usually won't perform much worse than GCN, but shines on graphs where "a few important neighbors determine the result"; the cost is more parameters and higher overfitting risk, needing stronger regularization on small graphs. Attention visualization can also give clues about "which nodes the model thinks are most critical" — but be careful: attention weights ≠ causal explanation, consistent with the warning in Transformer and NLP that "attention ≠ explanation."
7. Graph Pooling and Graph-Level Tasks
The previous sections discussed node-level representations. Graph classification (e.g., "is this molecule toxic?") requires turning the entire graph into a vector. There are two routes: global readout and hierarchical pooling.
1. Global readout (Readout)
The most naive approach is to perform a permutation-invariant aggregation over all node representations:
h_graph = READOUT( { h_v : v ∈ V } ) — mean / sum / max, or concatenation of all three"Average all nodes" seems simple but has a profound trap: averaging loses structural info. "Two rings" and "one four-ring" may have identical average features. This leads to the question of aggregation function choice: sum preserves the most info, mean is next, max is least. As we'll see in the expressiveness section, this choice directly determines the model's ceiling on graph isomorphism problems.
2. Hierarchical pooling: compressing the graph into "graphs"
A more powerful approach is hierarchical pooling: like CNN's spatial downsampling, shrink the graph layer by layer, preserving key structure. A representative is DiffPool (Ying et al., NeurIPS 2018): learn an assignment matrix that soft-assigns nodes into several "clusters," then treat clusters as new nodes and inter-cluster edges as new edges, shrinking the graph layer by layer. TopKPool takes a hard selection: at each layer, keep only the most important nodes by a learnable score.
Original graph (n nodes) → [GNN + pooling] → Compressed graph (k nodes) → [GNN + pooling] → … → readout → classifyThe benefit of hierarchical pooling is that "coarse-grained structures" (e.g., functional groups in molecules, communities in graphs) are explicitly modeled, often outperforming single readout on molecular and protein datasets.
3. Expressiveness ceiling: GIN and the WL test
Can a graph neural network "recognize" whether two graphs are isomorphic (same structure, just with node labels permuted)? This question has a beautiful theoretical answer: the expressive power limit of GNNs equals that of the Weisfeiler-Lehman (WL) graph isomorphism test — a classic algorithm that determines graph isomorphism by iteratively "labeling nodes and compressing the multiset of neighbor labels."
WL test per round:
① Each node collects its neighbors' current labels
② Hash the neighbor labels as a multiset, together with its own label, into a new label
③ If two graphs have different label sets in any round → graphs are not isomorphic
GIN per layer:
h_v^(k) = MLP^(k)( (1+ε) · h_v^(k-1) + Σ_{u∈N(v)} h_u^(k-1) )How Powerful are Graph Neural Networks? (Xu et al., ICLR 2019) proves: as long as the aggregation function is "injective" (sum satisfies this), GNN can achieve the WL test's expressive power; while mean/max lose multiset info (can't distinguish "two identical values" from "one value"), resulting in weaker expressiveness. This gives two engineering conclusions:
- sum aggregation has the strongest expressiveness (GIN therefore specifies sum aggregation + MLP update);
- no GNN can exceed 1-WL — exceeding it requires higher-order models (e.g., messages carrying higher-dimensional structures, equivariant models), with complexity increasing drastically.
8. Application Scenarios
GNN's value is ultimately demonstrated in how it lets deep learning consciously exploit relational structure for the first time. Four main battlefields:
1. Social networks
Social platforms are natural graphs: users are nodes, follow/friend relationships are edges. Typical tasks include:
- Friend recommendation: essentially link prediction — if two users share many common neighbors and those neighbors are familiar with each other (strong triadic closure), they're more likely to form a connection in the future;
- Community detection: cluster users' representations to find interest groups;
- Bot detection: social bots often manifest as "following many but getting zero follow-back, clustering around a few hubs" — such structural anomalies are hard for feature-based methods to capture.
2. Recommender systems
The core of recommender systems is the "user–item" bipartite graph: users and items are two types of nodes, and purchases/clicks/ratings are edges. GNN's value here is learning representations of users and items on the interaction graph, turning "collaborative filtering" into end-to-end graph learning:
- PinSage (KDD 2018): Pinterest's industrial solution, doing random walk sampling + aggregation on a 3-billion-edge graph for "pins and boards" recommendation, proven to significantly improve cold start and personalization;
- LightGCN: remove GCN's feature transformation and nonlinearity, keeping only aggregation — simple and performs excellently on mainstream recommendation datasets, another "less is more" case in GNN.
The broader landscape of recommender systems (including traditional collaborative filtering, two-tower models, and the ranking funnel) is in Recommender Systems; GNN is an important branch of "graph modeling" within it.
3. Molecular property prediction
Molecules are graphs of "atoms as nodes, chemical bonds as edges," and atom types and bond types are naturally discrete attributes — making GNNs the workhorse tool in quantum chemistry and drug discovery:
- Molecular fingerprints: traditional ECFP fingerprints hand-enumerate substructures; GNN uses message passing to automatically learn substructure representations, differentiable and end-to-end optimizable;
- Property regression: predicting quantum properties like HOMO-LUMO gaps and free energies on the QM9 dataset (130K small molecules), models like MPNN (the paper that proposed the message passing framework) significantly outperform hand-crafted features;
- Drug discovery: predicting toxicity, blood-brain barrier permeability, and drug–target interactions (essentially link prediction again). GNNs can also combine with large language models, fusing molecular SMILES strings with graph representations — a current research hotspot.
4. Knowledge graphs
Knowledge graphs represent the world as "entity–relation–entity" triplets (e.g., <Einstein, born in, Germany>), essentially heterogeneous graphs: nodes have types, edges have types. GNN must solve two problems here:
- Graph embeddings: early methods (TransE, Bordes et al. 2013) embed entities and relations in vector space, modeling relations via vector operations (
h + r ≈ t), but they can't exploit the graph's multi-hop structure; - Relational GCN (R-GCN, ESWC 2018): let each relation type have an independent transformation matrix, aggregate messages by relation type, generalizing GCN to heterogeneous graphs for graph completion (link prediction) and entity classification;
- Graph QA: parse questions into subgraphs, use GNNs to reason and find answer entities.
Other applications include: fraud detection (catching anomalous gangs in transfer graphs), traffic prediction (temporal GNNs on road networks), physics simulation (treating particles as nodes to predict trajectories), and program analysis (defect detection on code abstract syntax trees).
9. Limitations and Challenges
GNNs are far from omnipotent. Understanding its three fatal flaws is essential for engineers and researchers alike.
1. Over-smoothing
After deeper layers, node representations converge: repeatedly applying the normalized adjacency matrix is equivalent to repeatedly doing "neighbor averaging," and node features converge to the graph's stationary distribution as layers increase. Ultimately all nodes become nearly identical, and classification accuracy actually declines as layers increase — in stark contrast to CNN, where deeper means stronger.
Test accuracy
↑ /
│ / Typical GNN curve: accuracy drops when layers are too deep
│╱
├───────→ layers
1 2 3 4 5 6Mitigation methods include: residual connections (add input to output), JK-Net (concatenate outputs from all layers and learn weights), DropEdge (randomly drop edges during training), APPNP (replace neighbor aggregation with a tunable Personalized PageRank-style propagation). Engineering rule of thumb: GNNs rarely need more than 3–5 layers, unless the structure specifically handles over-smoothing.
2. Expressiveness ceiling
As discussed in Section 7, most GNNs' expressiveness doesn't exceed 1-WL graph isomorphism testing. This means:
- Certain structurally different graphs can never be distinguished by the model (e.g., "hexagonal ring" vs. "two triangular rings" may have identical representations under certain aggregations);
- The model fails when distinguishing complex substructures. Exceeding the limit requires higher-order GNNs (e.g., 2-WL, substructure-based models), but the computational cost is usually exponential, rarely used in engineering.
3. Scalability
Full-batch GCN needs the entire adjacency matrix and all intermediate representations in memory — millions of nodes approach a single GPU's limit; billions (real social networks) are directly infeasible. Common engineering routes:
| Strategy | Approach | Cost |
|---|---|---|
| Neighbor sampling | Take only a fixed-size subgraph per step (GraphSAGE approach) | May lose global structural info |
| Subgraph partitioning | Cluster-GCN and similar methods split the graph by community | Boundary info needs handling |
| Graph distillation/simplification | Learn a small graph to represent a large one | Lossy information |
Additionally, real graphs are often dynamic (edges add/remove over time) and heterogeneous (multi-type nodes and edges) — naively plugging static homophilic GNNs in usually doesn't work; dedicated design is required. These "data realities" have more first-hand resources in Datasets and Tools.
4. Over-squashing and long-range dependencies
The essence of message passing is "hop-by-hop diffusion," with information compressed at each hop. When a task needs to correlate two nodes far apart on the graph, distant information is often diluted after multiple hops — this is called over-squashing. It's isomorphic to RNN's vanishing gradient problem and is one of the root causes of "GNNs struggle with long-range reasoning." Mitigation directions include graph reconnection (reducing bottleneck edges) and learnable propagation distances.
Three easiest pitfalls on graphs
- Data leakage: neighbor normalization using all nodes, or any form of full-graph preprocessing before splitting the training set, artificially inflates test metrics — graph splits must be "split first, then compute structural info";
- Homophily assumption: GNN's naive design assumes "neighbors are more like myself"; naively applying GCN to heterophilic graphs (fraud, protein interactions, some e-commerce graphs) often underperforms linear baselines; first do a homophily coefficient test;
- Only comparing SOTA, not baselines: don't rush to the latest model — first run GCN, one-layer GraphSAGE, and a feature MLP baseline, confirming "structure truly helps" before talking about improvements. This is exactly the same first step as any machine learning project.
10. Tool Overview
Today, virtually no team writes message passing from scratch — two mature frameworks provide industrial-grade implementations and many built-in datasets.
1. PyTorch Geometric (PyG)
PyG is the de facto standard in the PyTorch ecosystem. Its core abstraction is torch_geometric.data.Data (encapsulating x, edge_index, y) and the MessagePassing base class:
python
import torch
import torch.nn.functional as F
from torch_geometric.nn import GCNConv
from torch_geometric.datasets import Planetoid
dataset = Planetoid(root="/tmp/data", name="Cora") # Built-in classic datasets
data = dataset[0] # data.x / edge_index / y
class Net(torch.nn.Module):
def __init__(self):
super().__init__()
self.conv1 = GCNConv(dataset.num_features, 16)
self.conv2 = GCNConv(16, dataset.num_classes)
def forward(self, x, edge_index):
x = F.relu(self.conv1(x, edge_index))
x = F.dropout(x, training=self.training)
return self.conv2(x, edge_index)
model = Net()
out = model(data.x, data.edge_index) # Semi-supervised training: only compute loss on labeled nodesPyG has built-in GCNConv, SAGEConv, GATConv, GINConv, TopKPooling, DiffPool, etc., and provides a full suite of standard datasets and evaluation protocols (Cora, QM9, OGB) — the OGB benchmark is also a public arena for evaluating GNNs.
2. Deep Graph Library (DGL)
DGL is known for its "graph as a first-class-citizen operator abstraction," decomposing message passing into send (construct messages) and recv (aggregate messages) two stages, supporting PyTorch/TensorFlow/MXNet backends:
python
import dgl
import torch.nn.functional as F
def gcn_message(edges):
return {"m": edges.src["h"]} # Message = source node features
def gcn_reduce(nodes):
return {"h": nodes.mailbox["m"].mean(dim=1)} # Aggregation = take mean
class GCNLayer(torch.nn.Module):
def forward(self, g, h, W):
g.ndata["h"] = h
g.update_all(gcn_message, gcn_reduce) # Built-in message passing scheduling
return F.relu(W(g.ndata["h"]))DGL's performance optimization (especially for large-scale sparse aggregation) has an excellent reputation in industry, and it also has high-level modules like dgl.nn.GraphConv and dgl.nn.GATConv. The choice between the two is more about ecosystem preference: deep PyTorch integration → PyG; cross-framework and ultra-large-scale aggregation performance → DGL.
3. Peripheral ecosystem
Graph preprocessing/visualization often uses NetworkX; molecule-specific libraries like RDKit (molecular graph construction) and DeepChem (drug chemistry ML toolkit); model and paper reproduction can reference Papers with Code's Graph leaderboard and the OGB benchmark. These resources are further organized in Datasets and Tools.
11. Tradeoffs and Decision Points
In practice, doing GNN is all choices. Here are four of the most core tradeoffs:
- Spectral vs. spatial methods: spectral methods like GCN are mathematically elegant and efficient, but normalized weights depend on the full graph structure, making them hard to transfer to dynamic graphs; spatial methods (GraphSAGE/GAT) depend only on "edge list + aggregation function," are flexible and inductive, and are today's mainstream.
- Mean vs. sum aggregation: mean is stable but weaker in expressiveness, can't distinguish multisets; sum has the strongest expressiveness (the basis of GIN) but can be amplified by high-degree nodes/outliers. Use mean for homogeneous graphs with uniform degree distributions; use sum when fine-grained structure distinction is needed.
- Layers vs. over-smoothing: 1–2 layer GNNs only encode local info; stacking to 5+ layers causes accuracy drops from over-smoothing. Determine the task's "dependency radius" before deciding layers; it's always better than blindly deepening; when needed, use residual connections, JK-Net, etc., to delay over-smoothing.
- Full graph vs. sampling: small graphs (up to tens of thousands of nodes) are easiest and most informative with full-batch; large graphs must sample or partition subgraphs, trading memory for structural info loss and training variance.
- Generic GNN vs. domain-specific: for homogeneous graphs, GCN/GAT work quickly out of the box; for heterogeneous graphs (knowledge graphs), dynamic graphs, or graphs with temporal aspects (traffic prediction), dedicated models are needed (R-GCN, temporal GNNs). "generic first, customize later" is the safest path.
One final principle, identical to other areas of machine learning: build simple baselines first (GCN, one-layer GraphSAGE) before talking about complex models. GNN paper SOTAs often rest on carefully tuned hyperparameters, and your business data may not eat that up. Understanding the evaluation of supervised learning and the training mechanics of deep learning, GNN is just a different data structure — the pitfalls (overfitting, data leakage, evaluation protocols) are exactly the same as textbooks describe. For quick term lookup (message passing, over-smoothing, WL test, Laplacian matrix), consult the Glossary anytime.
Further Reading
- What is Machine Learning — The discipline framework GNN belongs to and the four-element definition
- Deep Learning Foundations — GNN's neural network foundation: weight sharing, backpropagation
- Supervised Learning — The supervised paradigm behind node classification and link prediction
- Unsupervised Learning — The theoretical home of GraphSAGE's self-supervised loss
- Recommender Systems — The full recommendation landscape where PinSage / LightGCN live
- Transformer and NLP — Where GAT's attention mechanism comes from
- Large Language Models — Frontier directions for GNN + LLM fusion
- Glossary — Graph theory and GNN term quick reference
- Datasets and Tools — Cora, QM9, OGB, and PyG/DGL entry points
References
- Kipf & Welling. Semi-Supervised Classification with Graph Convolutional Networks (ICLR 2017) — Original GCN paper, the classic semi-supervised node classification setting
- Hamilton, Ying & Leskovec. Inductive Representation Learning on Large Graphs (NeurIPS 2017) — GraphSAGE: neighbor sampling and inductive learning
- Veličković et al. Graph Attention Networks (ICLR 2018) — GAT: one step of attention on graphs
- Gilmer et al. Neural Message Passing for Quantum Chemistry (ICML 2017) — MPNN: the unified message passing framework and molecular property prediction
- Defferrard, Bresson & Vandergheynst. Convolutional Neural Networks on Graphs with Fast Localized Spectral Filtering (NeurIPS 2016) — ChebNet: the key step making spectral methods practical
- Xu et al. How Powerful are Graph Neural Networks? (ICLR 2019) — GIN and the WL test: theoretical analysis of GNN expressiveness limits
- Li, Han & Wu. Deeper Insights into Graph Convolutional Networks for Semi-Supervised Learning (AAAI 2018) — Early systematic study of the over-smoothing phenomenon
- Xu et al. Representation Learning on Graphs with Jumping Knowledge Networks (ICML 2018) — JK-Net: adaptive multi-hop information combination
- Ying et al. Hierarchical Graph Representation Learning with Differentiable Pooling (NeurIPS 2018) — DiffPool: differentiable hierarchical pooling
- Ying et al. Graph Convolutional Neural Networks for Web-Scale Recommender Systems (KDD 2018) — PinSage: industrial-scale graph-based recommendation
- Schlichtkrull et al. Modeling Relational Data with Graph Convolutional Networks (ESWC 2018) — R-GCN: relational graph convolution on heterogeneous graphs (knowledge graphs)
- Bordes et al. Translating Embeddings for Modeling Multi-relational Data (NeurIPS 2013) — TransE: the foundational work of knowledge graph embeddings
- Wu et al. A Comprehensive Survey on Graph Neural Networks (IEEE TNNLS 2021) — Systematic GNN survey (Chinese and English summaries and categorization)
- PyTorch Geometric official documentation — PyG homepage, tutorials, and built-in datasets
- Deep Graph Library official documentation — DGL homepage and message passing API reference