Graph neural networks (GNNs)
Be able to explain message passing and train a GNN on a small graph.
Prerequisites
Intuition
A graph neural network works on data where the relations matter as much as the objects: molecules, social networks, road networks, citation graphs — and knowledge graphs like this platform's.
Message passing is the whole idea, and it is repeated in every layer:
1. Every node sends a message to its neighbours
2. Every node aggregates the messages it received (sum, mean, max)
3. Every node updates its representation from the aggregate and its own state
After layers every node has information from every node within steps. That is why depth in a GNN means something entirely different from depth in a CNN: it governs how far information travels in the graph, not how abstract the features become.
The aggregation has to be permutation-invariant — the neighbours have no order. Hence a sum, a mean or a max, never a concatenation.
Formal
Three task types, and they require different splits:
| Type | Example | Split |
|---|---|---|
| Node level | classify users in a network | mask nodes in the same graph |
| Edge level | predict links | mask edges |
| Graph level | a molecule's property | split graphs between the sets |
The node-level split is special: the whole graph is visible during training, but only certain nodes' labels are used. That is called transductive learning and means that ordinary leakage checks are not enough — the neighbours' features are visible for the test nodes too.
Two problems that make deep GNNs hard:
| Problem | What happens |
|---|---|
| Over-smoothing | after many layers all the node representations converge towards the same vector |
| Over-squashing | information from exponentially many nodes is pressed through a narrow bottleneck |
That is why most GNNs are shallow — 2 to 4 layers. Going deeper requires residual connections, normalisation and sometimes graph rewiring.
Architecture families:
| Model | Aggregation |
|---|---|
| GCN | a normalised mean over the neighbours |
| GraphSAGE | sample a fixed number of neighbours — scalable to large graphs |
| GAT | attention over the neighbours, learnt weights |
| GIN | a sum plus an MLP; proved as expressive as the Weisfeiler–Lehman test |
The GIN result is theoretically important: it sets an upper bound on what message-passing GNNs can distinguish. Two graphs the WL test cannot tell apart cannot be told apart by a standard GNN either — certain regular graphs, for instance.
Practical advice:
- Start with a baseline without the graph. An MLP on the nodes' own features, or logistic regression. If the GNN does not beat it, the graph adds nothing.
- Use 2–3 layers.
- Add the node degree and other structural features — often a larger effect than a finer architecture.
- For large graphs: neighbour sampling (GraphSAGE) or cluster-based batching.
Advice 1 is the one most often skipped, and it decides whether the project is meaningful.
Code
import torch, torch.nn as nn, torch.nn.functional as F
class GCNLayer(nn.Module):
"""h' = σ(D^-1/2 Â D^-1/2 h W), Â = A + I"""
def __init__(self, cin, cout):
super().__init__()
self.lin = nn.Linear(cin, cout)
def forward(self, h, A_norm):
return A_norm @ self.lin(h)
def normalise(A):
A_hat = A + torch.eye(len(A)) # add self-loops
degree = A_hat.sum(1)
d = torch.diag(degree.pow(-0.5))
return d @ A_hat @ d
class GCN(nn.Module):
def __init__(self, in_dim, hidden, n_classes, layers=2, dropout=0.5):
super().__init__()
dims = [in_dim] + [hidden] * (layers - 1) + [n_classes]
self.layers = nn.ModuleList([GCNLayer(dims[i], dims[i + 1]) for i in range(layers)])
self.dropout = dropout
def forward(self, h, A_norm):
for i, lg in enumerate(self.layers):
h = lg(h, A_norm)
if i < len(self.layers) - 1:
h = F.dropout(F.relu(h), self.dropout, self.training)
return h
# A small graph: four nodes, a square
A = torch.tensor([[0., 1., 0., 1.],
[1., 0., 1., 0.],
[0., 1., 0., 1.],
[1., 0., 1., 0.]])
An = normalise(A)
h = torch.randn(4, 8)
model = GCN(8, 16, 2)
print(model(h, An).shape) # torch.Size([4, 2])
# Over-smoothing: measure how similar the representations become with depth
def similarity(h):
hn = F.normalize(h, dim=-1)
s = hn @ hn.T
n = len(h)
return float((s.sum() - s.diag().sum()) / (n * (n - 1)))
h = torch.randn(4, 8)
for depth in (1, 2, 4, 8, 16):
m = GCN(8, 8, 8, layers=depth)
m.eval()
with torch.no_grad():
print(f" {depth:>2} layers: average cosine similarity {similarity(m(h, An)):.4f}")
# 1 layer: 0.1...
# 16 layers: near 1.0 ← all the nodes have become alike: over-smoothing
# Transductive node classification: mask the LABELS, not the nodes
def split_nodes(n, train_share=0.6, val_share=0.2, seed=0):
g = torch.Generator().manual_seed(seed)
perm = torch.randperm(n, generator=g)
a, b = int(n * train_share), int(n * (train_share + val_share))
mask = lambda idx: torch.zeros(n, dtype=torch.bool).index_fill_(0, idx, True)
return mask(perm[:a]), mask(perm[a:b]), mask(perm[b:])
# A baseline WITHOUT the graph — always run this first
class MLPBaseline(nn.Module):
def __init__(self, in_dim, hidden, n_classes):
super().__init__()
self.f = nn.Sequential(nn.Linear(in_dim, hidden), nn.ReLU(),
nn.Dropout(0.5), nn.Linear(hidden, n_classes))
def forward(self, h, A_norm=None):
return self.f(h)
# If the GCN does not beat this, the graph structure adds nothing to the task.
Mastery means
- Explains message passing
- Knows why deep GNNs are hard
- Chooses the task type and the split correctly
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Semi-Supervised Classification with Graph Convolutional Networks — arXiv (open access; licence per article)
- arXiv — How Powerful are Graph Neural Networks? (GIN) — arXiv (open access; licence per article)
- PyTorch Geometric — dokumentation (MIT) — MIT