🔗 Deep Learning · Lecture 33 of 38

Graph Neural Networks: Learning on Relational Data

Molecules, social networks, road maps and knowledge graphs are graphs. GNNs learn from them by passing messages between neighbours. We derive message passing, GCN and GAT layers, and survey node, edge and graph-level tasks.

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:

  1. Collects messages from its neighbours;
  2. Aggregates them with a permutation-invariant function (sum, mean, max);
  3. Updates its own representation.
$$ \mathbf{h}_v^{(l+1)} = \text{UPDATE}\left(\mathbf{h}_v^{(l)},\; \text{AGG}_{u \in \mathcal{N}(v)}\,\text{MSG}\big(\mathbf{h}_v^{(l)}, \mathbf{h}_u^{(l)}, \mathbf{e}_{uv}\big)\right) $$

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}$):

$$ \mathbf{H}^{(l+1)} = \sigma\left(\tilde{\mathbf{D}}^{-1/2}\tilde{\mathbf{A}}\tilde{\mathbf{D}}^{-1/2}\mathbf{H}^{(l)}\mathbf{W}^{(l)}\right) $$

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:

$$ \alpha_{vu} = \text{softmax}_{u \in \mathcal{N}(v)}\Big(\text{LeakyReLU}\big(\mathbf{a}^\top[\mathbf{W}\mathbf{h}_v \,\|\, \mathbf{W}\mathbf{h}_u]\big)\Big), \qquad \mathbf{h}_v' = \sigma\Big(\sum_{u \in \mathcal{N}(v)}\alpha_{vu}\mathbf{W}\mathbf{h}_u\Big) $$

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#

python
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#

LevelExampleReadout
NodeClassify users, papers, proteinsNode embeddings
EdgePredict links: friendships, drug–target interactions, recommendationsPairs of node embeddings
GraphPredict molecular toxicity or solubilityPool 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.

JA
Written by

Janin A Apurba

B.Sc. in CSE, AUST · Advanced ICT Officer, CNRS-UNHCR. Teaching AI, ML and Deep Learning to the next generation of engineers and researchers.

Keep learning

Related lectures

🔗 Deep Learning

Distributed Training: Data, Model, Pipeline and Sharded Parallelism

Large models and datasets need many accelerators. We explain data parallelism with all-reduce, sharded data parallelism (ZeRO/FSDP), tensor and pipeline model parallelism, and how they combine at scale.

Advanced⏱ 6 min#128
🔗 Deep Learning

Knowledge Distillation: Teaching Small Models with Large Ones

A large "teacher" model's soft predictions contain rich information that can train a much smaller "student". We derive the distillation loss with temperature, discuss dark knowledge, and survey feature and LLM distillation.

Intermediate⏱ 5 min#130
🔗 Deep Learning

Mixed-Precision Training and GPU Efficiency

Training in 16-bit arithmetic roughly halves memory and can multiply throughput. We explain float16 and bfloat16, loss scaling, automatic mixed precision, and other practical techniques to make GPUs work harder.

Advanced⏱ 5 min#127