Skip to content

Graph Neural Networks

Quick overview Graph neural networks extend deep learning to structured data on nodes and edges. This article unpacks graph representations, the message passing framework, and three mainstream models — GCN, GAT, and GraphSAGE — covering node classification, link prediction, and graph classification tasks and applications (recommenders, molecules, physics simulation), while confronting limitations like neighbor sampling, over-smoothing, and scalability.

Graph Neural Networks ​

In a sentence: Graph Neural Networks (GNNs) learn graph structure and node representations through repeated "message passing" between nodes — each node updates itself using neighbor information, iteratively layering to encode representations into vectors — they extend deep learning's territory from "regular grids/sequences" to "irregular relational data" (social networks, molecules, knowledge graphs, transaction chains).

1. Graph Data Representation ​

A graph $G = (V, E)$ consists of nodes $V$ and edges $E$. To feed graphs into deep learning, we first need to convert them to tensors, with common representations:

  • Adjacency matrix A: An $n \times n$ 0/1 matrix (or weighted), where $A[i][j] = 1$ means there is an edge $i \rightarrow j$. Problem: quadratic scale, permutation-sensitive, sparse;
  • Feature matrix X: Attributes of each node/edge (user age, item category, molecular atom type);
  • Homogeneous vs. heterogeneous graphs: Single node/edge type vs. diverse types (e.g., three node types — users, items, stores — in recommendation systems);
  • Adjacency lists / sampled views: The practical representation for large industrial graphs (sparse storage).

The GNN input-output paradigm is the same as standard networks — input node features and structure, output node/graph/edge vector representations. This is Representation Learning and Pretraining concretized on graph data.

2. Message Passing Framework ​

Most GNNs share the same template — Message Passing (the unified formulation proposed by Gilmer et al. 2017), performing two steps per layer:

  1. Aggregate: Collect neighbor representations, via sum, mean, max, attention-weighted, etc.;
  2. Update: Combine the aggregated result with the node's own representation through a nonlinear transformation to get a new representation: $$ h_v^{(k+1)} = \text{UPDATE}\left(h_v^{(k)},; \text{AGG}\left({h_u^{(k)} : u \in \mathcal{N}(v)}\right)\right) $$

After $K$ layers, a node's representation contains information from its $K$-hop neighbors — the receptive field expands with depth, directly corresponding to the CNN receptive field concept (see CNNs and Computer Vision). The value of the message passing framework is that it unifies the "aggregate/update" choices across various GNNs, and exposes their shared limitations (see over-smoothing below).

3. Three Mainstream Models ​

ModelYearAggregationCharacteristics
GCN2017Neighbor mean (with degree normalization)Simple and effective, spectral graph theory's convolution approximation; cannot distinguish neighbor importance
GraphSAGE2017Sample fixed-number neighbors + aggregate (mean/max/LSTM)Supports large scale and inductive learning (new nodes without retraining)
GAT2018Attention-weighted aggregationNeighbor importance is learnable, higher theoretical ceiling
  • GCN (Graph Convolutional Network): $H^{(l+1)} = \sigma(\hat D^{-1/2}\hat A \hat D^{-1/2} H^{(l)} W^{(l)})$, which is "convolution" on graphs — but understand it essentially as normalized neighbor message weighting, not the sliding window of image convolution;
  • GraphSAGE: Rather than stuffing the entire graph into the network, it samples neighbors per node, decoupling training from graph scale — the first step toward industrial deployment;
  • GAT (Graph Attention Network): Uses ideas from attention mechanisms to learn weights for each edge, "weighting neighbors appropriately" during aggregation — nodes with many neighbors won't drown out key information.

4. Graph Tasks ​

  • Node classification: Label unlabeled nodes (e.g., "community/identity" recognition in social networks). Common in semi-supervised settings — label only some nodes, learn globally;
  • Link prediction: Determine whether two nodes should be connected (friend recommendations, knowledge graph completion, molecular bond prediction);
  • Graph classification/regression: A single label for the entire graph (whether a molecule is active, whether a program compiles);
  • Graph generation: Generate new molecules, new network structures (combined with Generative Models).

5. Applications ​

  • Recommendation systems: PinSage uses graph convolution to aggregate "item neighbors" for recall — the benchmark for GNN industrial deployment (full pipeline in Deep Learning Recommender Systems);
  • Molecular property prediction: Molecules = atom nodes + chemical bond edges; GNNs learn molecular representations for drug screening, material property regression (QM9 benchmarks, etc.);
  • Physics simulation: Model fluids/particles/cloth as graphs (nodes = particles, edges = interactions); GNNs simulate dynamical evolution (GNS, etc.);
  • Knowledge graphs: Entity-relation completion, question-answering search, combined with RAG to give Large Language Models (LLM) structured knowledge;
  • Code analysis: Parse programs into abstract syntax trees (ASTs, a type of directed graph) for vulnerability detection, code completion.

6. Scale and Sampling ​

GNNs' industrial bottleneck is graphs being too large (billions of nodes): whole-graph training is infeasible in memory and propagation cost. Main countermeasures:

  1. Neighbor sampling (invented by GraphSAGE): Each node samples a fixed number of neighbors (e.g., 10) per layer, turning "whole-graph propagation" into "mini-batch subgraph propagation";
  2. Subgraph sampling: Split training by subgraphs (Cluster-GCN, GraphSAINT), balancing class distributions;
  3. Distributed computing and hardware: Partition and store graph shards, sparse-matrix-specialized operators (e.g., DGL/GraphScope optimizations);
  4. Scalable architectures: SIGN/SGC precompute neighborhood aggregation, skipping per-layer propagation, trading accuracy for scale.

On the engineering framework side, PyTorch Geometric (PyG) and DGL are the mainstream choices (selection comparison in How to Choose Frameworks and Tools).

7. Limitations: Over-Smoothing and Scalability ​

  • Over-smoothing: As layers deepen, node representations converge to the same value (all become "neighbor averages"), losing discriminability — "gradient disappearance on graphs." Countermeasures: residual connections, skip connections, PairNorm, stochastic depth. This is the same root cause as "why deep networks are hard to train" in Initialization and Normalization;
  • Heterogeneity challenges: Real-world graphs are mostly heterogeneous (different node/edge types); standard GNNs aggregate uniformly, losing type information — requiring specialized designs like HAN/RGCN;
  • Dynamic graphs: Edges change in real time (social networks, transactions); statically trained GNNs can't keep up — requiring incremental/temporal models;
  • Scalability ceiling: Even with sampling, tens of billions of nodes + high-cardinality features remain an engineering challenge;
  • Expressivity ceiling: Standard message-passing GNNs have expressivity no better than the 1-WL (Weisfeiler-Lehman) graph isomorphism test, unable to distinguish certain structures (GTN and more complex structures attempt to break through).

Practical Advice

In industrial applications, graph features (degree, clustering coefficient, PageRank) combined with a simple MLP often outperform complex GNNs initially. GNN benefits mostly come from joint encoding of "structure + attributes." Build a baseline first, then add structural models — this is a universal engineering principle (see Training Recipes and Hyperparameter Tuning).

8. Trade-offs ​

  • Expressivity vs complexity: Attention aggregation (GAT) is strong but expensive; mean aggregation (GCN) is cheap but blunt; for resource-constrained scenarios, start with GCN/GraphSAGE baselines;
  • Depth vs over-smoothing: Most real tasks need only 2–3 layers (the receptive field already covers "friends of friends"); blindly adding layers degrades performance;
  • Inductive vs transductive: Choose inductive (GraphSAGE-style) for "never-before-seen new nodes/graphs"; transductive (full-graph Laplacian) suffices for serving a fixed graph;
  • GNN vs Transformer: Graph attention and Transformers share the same roots (graphs are a special case of "sparse attention"); large graphs benefit from GNN sampling, while small graphs/strongly structured data often use GTN-style approaches — the two are converging, see Transformer Architecture.

Further Reading ​

References ​