Skip to content
AI-grafen
FAI engineeringDeep learning· about 90 min· fast-moving, sources checked often· verified 2026-09-21· EN

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

hv(l+1)=ϕ(hv(l), ⨁u∈N(v)ψ(hv(l),hu(l))) h_v^{(l+1)} = \phi\left(h_v^{(l)},\ \bigoplus_{u \in \mathcal{N}(v)} \psi\left(h_v^{(l)}, h_u^{(l)}\right)\right)\

After kk layers every node has information from every node within kk 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:

TypeExampleSplit
Node levelclassify users in a networkmask nodes in the same graph
Edge levelpredict linksmask edges
Graph levela molecule's propertysplit 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:

ProblemWhat happens
Over-smoothingafter many layers all the node representations converge towards the same vector
Over-squashinginformation 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:

ModelAggregation
GCNa normalised mean over the neighbours
GraphSAGEsample a fixed number of neighbours — scalable to large graphs
GATattention over the neighbours, learnt weights
GINa 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:

  1. 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.
  2. Use 2–3 layers.
  3. Add the node degree and other structural features — often a larger effect than a finer architecture.
  4. 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

All the sources and licences