Images live on grids and text on sequences, but much of the world's data is relational: atoms bonded in molecules, people connected in social networks, roads linking towns, citations between papers, transactions between accounts. Such data forms a graph $G = (V, E)$ — nodes with features, connected by edges. Graphs have no fixed ordering or size, so CNNs and RNNs do not apply directly. Graph Neural Networks (GNNs) extend deep learning to this setting.
Graph basics#
- Adjacency matrix $\mathbf{A} \in \{0,1\}^{n \times n}$: $A_{ij} = 1$ if nodes $i$ and $j$ are connected.
- Degree matrix $\mathbf{D}$: diagonal, $D_{ii} = \sum_j A_{ij}$.
- Node features $\mathbf{X} \in \mathbb{R}^{n \times d}$; optionally edge features.
A key requirement: the output should not depend on the arbitrary order in which we number the nodes. Node-level outputs must be permutation equivariant; graph-level outputs must be permutation invariant.
Message passing: the unifying framework#
Almost all GNNs follow the message-passing paradigm (Gilmer et al., 2017). In each layer, every node:
- Collects messages from its neighbours;
- Aggregates them with a permutation-invariant function (sum, mean, max);
- Updates its own representation.
After $k$ layers, each node's representation depends on its $k$-hop neighbourhood — analogous to a CNN's growing receptive field.
Graph Convolutional Network (GCN)#
Kipf and Welling's GCN (2017) uses a normalised average of neighbour features, including a self-loop ($\tilde{\mathbf{A}} = \mathbf{A} + \mathbf{I}$):
The symmetric normalisation prevents high-degree nodes from dominating and keeps feature scales stable. The weight matrix $\mathbf{W}^{(l)}$ is shared across all nodes — weight sharing, as in convolution.
Graph Attention Network (GAT)#
GCN weights neighbours by fixed degree-based coefficients. GAT (Veličković et al., 2018) learns them with attention:
Different neighbours can matter differently, and multi-head attention stabilises learning. (A transformer is, in fact, a GAT on a fully connected graph of tokens.)
Other important layers#
- GraphSAGE — samples a fixed number of neighbours and aggregates them, enabling inductive learning on huge graphs (used for recommendation at web scale).
- GIN (Graph Isomorphism Network) — uses sum aggregation and an MLP; provably as expressive as the Weisfeiler–Lehman graph isomorphism test, the theoretical limit for standard message passing.
- Edge-conditioned / MPNN layers — use edge features (bond types, distances) in messages; common in chemistry.
- Equivariant GNNs — respect 3-D rotations and translations for molecules and physics.
A GCN from scratch#
import torch
import torch.nn as nn
import torch.nn.functional as F
class GCNLayer(nn.Module):
def __init__(self, d_in, d_out):
super().__init__()
self.lin = nn.Linear(d_in, d_out, bias=False)
def forward(self, X, A):
A_hat = A + torch.eye(A.size(0)) # add self-loops
d = A_hat.sum(1)
D_inv_sqrt = torch.diag(d.pow(-0.5))
return D_inv_sqrt @ A_hat @ D_inv_sqrt @ self.lin(X)
class GCN(nn.Module):
def __init__(self, d_in, d_hid, n_cls):
super().__init__()
self.l1, self.l2 = GCNLayer(d_in, d_hid), GCNLayer(d_hid, n_cls)
def forward(self, X, A):
return self.l2(F.dropout(F.relu(self.l1(X, A)), 0.5, self.training), A)
# Two communities of 20 nodes with dense internal links and a few cross links
torch.manual_seed(0)
n = 40; labels = torch.tensor([0] * 20 + [1] * 20)
A = (torch.rand(n, n) < 0.02).float()
same = labels[:, None] == labels[None, :]
A = torch.maximum(A, ((torch.rand(n, n) < 0.3) & same).float())
A = torch.triu(A, 1); A = A + A.T # symmetric, no self-loops
X = torch.randn(n, 8) # uninformative features
train_mask = torch.zeros(n, dtype=torch.bool); train_mask[[0, 1, 20, 21]] = True # 4 labels only
model = GCN(8, 16, 2); opt = torch.optim.Adam(model.parameters(), 0.01, weight_decay=5e-4)
for epoch in range(200):
model.train(); opt.zero_grad()
loss = F.cross_entropy(model(X, A)[train_mask], labels[train_mask]); loss.backward(); opt.step()
model.eval()
print("accuracy on all nodes:", (model(X, A).argmax(1) == labels).float().mean().item())Even with random node features and only four labelled nodes, the GCN classifies most nodes correctly by exploiting the graph structure — a semi-supervised setting where GNNs shine. In practice, use libraries such as PyTorch Geometric or DGL, which handle sparse adjacency efficiently.
Tasks#
| Level | Example | Readout |
|---|---|---|
| Node | Classify users, papers, proteins | Node embeddings |
| Edge | Predict links: friendships, drug–target interactions, recommendations | Pairs of node embeddings |
| Graph | Predict molecular toxicity or solubility | Pool all nodes (sum/mean/attention) |
Challenges#
- Over-smoothing — after many layers, node representations become indistinguishable, so GNNs are usually shallow (2–4 layers); residual connections and normalisation help.
- Over-squashing — information from exponentially growing neighbourhoods is compressed into fixed-size vectors, limiting long-range reasoning; graph rewiring and graph transformers address this.
- Expressivity limits — standard message passing cannot distinguish certain non-isomorphic graphs.
- Scalability — graphs with billions of edges need neighbour sampling and mini-batching.
Applications#
Drug discovery and molecular property prediction, materials science, traffic and travel-time forecasting (used in widely deployed map services), fraud detection in transaction networks, recommendation, knowledge-graph completion, physics simulation, and protein-structure pipelines where graph-like reasoning over residues plays a central role.