A graph neural network (GNN) learns from data whose structure is a set of entities and the links between them: users and items, atoms and bonds, accounts and transactions, papers and citations. Flattening such data into a table throws away the links. Feeding the adjacency matrix to an ordinary MLP ties the model to one arbitrary node ordering and one graph size. GNNs avoid both problems by building every layer from operations that do not care how nodes are numbered.
This article starts from that requirement, derives message passing, implements a GCN and a GraphSAGE layer in plain PyTorch, works a numeric example by hand, and then covers what goes wrong: limited expressive power, oversmoothing, neighbour explosion, and leaky evaluation. By the end you should be able to train a node classifier and know which numbers to watch.
From permutation symmetry to message passing
Write a graph as a node feature matrix X (n rows, d features) and an edge list. If you relabel the nodes with a permutation P, the graph is the same, so a node-level model f must satisfy f(PX, PAPT) = P f(X, A): its outputs permute along with its inputs (equivariance). A graph-level prediction must not change at all (invariance). The simplest way to guarantee this is to compute each node's new state from its own state and an order-independent aggregate, such as a sum, mean or max, of its neighbours' states, using the same weights at every node. That is message passing:
for layer k = 1..K:
for every edge (u -> v): m_uv = MESSAGE_k(h_u, h_v, e_uv)
for every node v: a_v = AGGREGATE_k({ m_uv : u in N(v) }) # sum / mean / max
h_v = UPDATE_k(h_v, a_v)
graph readout (optional): h_G = READOUT({ h_v }) # sum / mean poolingAfter K layers, node v's state depends on its K-hop neighbourhood, much as a K-step breadth-first search reaches K hops. Because the weights are shared, the model has the same number of parameters for a 10-node molecule and a 100-million-node social graph.
GCN: normalised neighbourhood averaging, worked by hand
The graph convolutional network (GCN) of Kipf and Welling is the most common baseline. One layer computes H' = σ( H W), where  = D-1/2(A + I)D-1/2. Adding I gives every node a self-loop so its own features survive. D is the degree matrix of A + I. The symmetric normalisation scales the message from u to v by 1/√(deg(u)·deg(v)), so high-degree nodes neither shout over nor drown out their neighbours, and repeated layers do not blow up the activations.
Worked example. Take a path 1-2-3 with self-loops. Degrees are 2, 3 and 2. The weights are 1/2 for the self-loops at nodes 1 and 3, 1/3 for node 2's self-loop, and 1/√6 ≈ 0.408 for each real edge. Give node 1 the feature 1 and the others 0, and leave out W and σ to watch only the propagation. After one layer the states are [0.5, 0.408, 0]: node 3 has heard nothing yet. After two layers they are [0.25 + 0.167, 0.204 + 0.136, 0.167] = [0.417, 0.340, 0.167]. The signal took two hops to reach node 3 and is already spreading evenly. Keep that evening-out in mind; it becomes oversmoothing.
In code, the sparse product ÂH is a gather along edges followed by a scatter-add into destinations. index_add_ does the scatter, so no graph library is needed:
import torch
import torch.nn as nn
import torch.nn.functional as F
def gcn_norm(edge_index, num_nodes):
"""Self-loops plus symmetric weights 1/sqrt(deg(u) deg(v)). edge_index: [2, E], both directions."""
loops = torch.arange(num_nodes, device=edge_index.device)
src = torch.cat([edge_index[0], loops])
dst = torch.cat([edge_index[1], loops])
deg = torch.zeros(num_nodes, device=src.device).index_add_(
0, dst, torch.ones(dst.numel(), device=src.device))
inv_sqrt = deg.pow(-0.5) # deg >= 1 thanks to self-loops
return src, dst, inv_sqrt[src] * inv_sqrt[dst]
class GCNLayer(nn.Module):
def __init__(self, d_in, d_out):
super().__init__()
self.lin = nn.Linear(d_in, d_out, bias=False)
self.bias = nn.Parameter(torch.zeros(d_out))
def forward(self, x, src, dst, w):
h = self.lin(x) # transform first: E * d_out work, not E * d_in
msg = h[src] * w.unsqueeze(-1) # one message per edge
return torch.zeros_like(h).index_add_(0, dst, msg) + self.bias
class GCN(nn.Module):
def __init__(self, d_in, d_hid, n_cls, p=0.5):
super().__init__()
self.l1, self.l2, self.p = GCNLayer(d_in, d_hid), GCNLayer(d_hid, n_cls), p
def forward(self, x, g):
x = F.dropout(x, self.p, self.training)
x = F.relu(self.l1(x, *g))
x = F.dropout(x, self.p, self.training)
return self.l2(x, *g)
Training a node classifier
Node classification on a single graph is usually transductive: every node and edge is visible during training, but only the training nodes' labels enter the loss. The loop is ordinary PyTorch with boolean masks. Make the edge list symmetric first, or messages flow one way only:
def train_gcn(x, edge_index, y, train_mask, val_mask, epochs=200):
edge_index = torch.unique(torch.cat([edge_index, edge_index.flip(0)], 1), dim=1)
g = gcn_norm(edge_index, x.size(0))
model = GCN(x.size(1), 64, int(y.max()) + 1)
opt = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
best_acc, best_state = -1.0, None
for epoch in range(epochs):
model.train()
opt.zero_grad()
loss = F.cross_entropy(model(x, g)[train_mask], y[train_mask])
loss.backward()
opt.step()
model.eval()
with torch.no_grad():
pred = model(x, g).argmax(-1)
acc = (pred[val_mask] == y[val_mask]).float().mean().item()
if acc > best_acc:
best_acc = acc
best_state = {k: v.clone() for k, v in model.state_dict().items()}
model.load_state_dict(best_state)
return model, best_accTwo layers is the usual default and is often the best choice. Deeper plain GCNs tend to get worse, not better, for reasons covered below.
Graph-level tasks, such as predicting a property of a whole molecule, need two changes. First, batch many small graphs as one disconnected graph: concatenate their feature matrices, add each graph's node offset to its edge indices, and keep a batch vector recording which graph owns each node. No edge crosses between graphs, so the layers above work unchanged. Second, add a readout after the last layer that sums or averages node states per graph, again with index_add_ keyed by that vector, then pass the result to a small MLP head. Sum readout preserves graph size information and mean readout discards it, so choose deliberately.
Serving. For node-level predictions that do not need to be fresh, run the trained model over the full graph offline and store each node's embedding. Online inference then becomes a lookup, and new edges are only reflected at the next refresh.
Scaling: GraphSAGE and neighbour sampling
Full-graph training needs every node's activations for every layer in memory. A graph with 100 million nodes, 256 hidden units and 3 layers needs about 300 GB of FP32 activations before gradients. Mini-batching is the fix, but it has a trap: a node's K-layer output depends on its whole K-hop neighbourhood, which in a social graph can be most of the graph (neighbour explosion).
GraphSAGE (Hamilton, Ying and Leskovec) addresses this with neighbour sampling. For each seed node it samples a fixed number of neighbours per layer, for example 15 and then 10. That caps one seed's computation at 1 + 15 + 150 nodes however large the hubs are. Because it learns aggregation functions over features rather than an embedding per node, it can embed nodes unseen during training (inductive learning). Its layer keeps separate weights for the node itself and its neighbourhood:
class SAGELayer(nn.Module):
def __init__(self, d_in, d_out):
super().__init__()
self.self_lin = nn.Linear(d_in, d_out)
self.nbr_lin = nn.Linear(d_in, d_out, bias=False)
def forward(self, x, edge_index):
src, dst = edge_index
n = x.size(0)
total = torch.zeros(n, x.size(1), device=x.device).index_add_(0, dst, x[src])
count = torch.zeros(n, device=x.device).index_add_(
0, dst, torch.ones(src.numel(), device=x.device))
mean = total / count.clamp(min=1).unsqueeze(-1) # isolated nodes keep 0
return self.self_lin(x) + self.nbr_lin(mean)In PyTorch Geometric the sampler is NeighborLoader(data, num_neighbors=[15, 10], batch_size=1024, input_nodes=data.train_mask). Each batch is a subgraph whose first batch.batch_size nodes are the seeds, so compute the loss on out[:batch.batch_size] only. Sampling makes training stochastic over structure as well as data, so for evaluation run the model layer by layer over the full graph, or average several sampled passes.
Two other aggregators are worth knowing. Graph attention networks (GAT) weight neighbours with learned attention scores, the same idea as attention in transformers restricted to graph edges. GIN (Xu et al.) uses a sum followed by an MLP, for a reason explained next.
Limits: expressivity, oversmoothing, oversquashing
Expressive power. Xu et al. proved that message-passing GNNs can tell two graphs apart no better than the 1-dimensional Weisfeiler-Lehman (1-WL) colour-refinement test, and that sum aggregation with an injective update (GIN) reaches that bound while mean and max do not. A concrete failure: two disjoint triangles and a single 6-cycle. Both have six nodes, each of degree 2. With identical input features every node receives the same messages at every layer, so no GNN of this family can tell them apart or count triangles. If your task depends on cycles or substructures, such as rings in molecules, add structural features (cycle counts, random-walk or Laplacian positional encodings) or use a higher-order model.
Oversmoothing. Each GCN layer is a weighted averaging step, and the worked example showed values evening out. After many layers node states in a connected component converge toward each other and become useless for classification. Residual connections, jumping-knowledge concatenation of all layers' outputs, normalisation, or simply fewer layers help.
Oversquashing. When information has to travel far, an exponentially growing neighbourhood is compressed into a fixed-width vector. Bottleneck edges, such as the bridge between two communities, choke it. Adding virtual nodes or rewiring the graph helps, and so does handing long-range parts of the problem to a graph transformer.
Failure modes in practice
- One-directional edges. An edge list containing only u→v means v hears from u but not the reverse. Symmetrise undirected graphs explicitly, and deduplicate, because a duplicated edge counts twice in a sum.
- Label leakage. Using labels, or features derived from labels, as node inputs lets validation labels flow to neighbours through message passing. In fraud graphs a "known fraudster neighbour" feature computed on the full label set is the classic case.
- Temporal leakage. Building the graph from all edges, including ones created after the prediction time, inflates offline accuracy. Split by time and build the training graph only from earlier edges.
- Random splits on clustered graphs. Neighbouring nodes are similar (homophily), so random node splits put near-duplicates on both sides. Report a split by community or time as well.
- Heterophily. When linked nodes tend to have different labels, as with fraudsters linked to victims, plain averaging hurts. Always compare against an MLP on node features alone; if the GNN does not beat it, the graph is not helping.
- Hubs and memory. A celebrity node with millions of edges dominates memory and latency. Cap the fan-in with sampling at both training and serving time.
Trade-offs
| Choice | Use when | Cost |
|---|---|---|
| GCN | Homophilous graph, fits in memory, quick baseline | Usually trained full-graph; oversmooths when deep |
| GraphSAGE + sampling | Large graphs, new nodes at serving time | Stochastic outputs; sampler tuning |
| GAT | Neighbours differ in relevance | More memory per edge; slower |
| GIN | Graph classification where structure matters | Still bounded by 1-WL |
| MLP on features | Always, as the baseline to beat | Ignores structure |
GNN layers are sparse gathers and scatters, so on accelerators they are bandwidth-bound and sensitive to edge ordering. Sorting edges by destination improves locality, and compiling the model (see torch.compile on GPUs) can fuse the elementwise work around them. Link-level ranking over graphs connects to PageRank, which is message passing with fixed rather than learned weights.
What to do next
- Train an MLP on node features alone and record its validation accuracy as the floor.
- Implement the two-layer GCN above, symmetrise and deduplicate the edges, and beat that floor.
- Check the degree distribution; if the largest degree is in the thousands, switch to GraphSAGE with a sampled fan-out such as [15, 10].
- Audit every input feature for label and temporal leakage, and add a time- or community-based split.
- Try 2, 3 and 4 layers with residual connections and plot accuracy against depth to see where oversmoothing starts.
- If the task depends on cycles or motifs, add structural or positional encodings before reaching for a bigger model.