Graph Neural Networks Explained
Learn how graph neural networks use message passing, which architecture to choose, how to train them, and where their limits matter.
Graph neural networks (GNNs) are neural models that learn from entities and the relationships between them. Instead of treating every example as an independent row, a GNN uses a graph’s nodes, edges, and optional features as part of the input. It repeatedly exchanges information between connected nodes, then uses the resulting representations for node, edge, or whole-graph predictions.
This makes GNNs useful for molecules, recommender systems, social networks, knowledge graphs, physical simulations, 3D data, and transaction networks. It also creates distinctive engineering problems: graph sampling, information leakage, oversmoothing, oversquashing, changing topology, and expensive memory use.
What is a graph neural network?
A graph is usually written as G = (V, E):
- Nodes (V): entities such as users, products, atoms, documents, or sensors.
- Edges (E): relationships such as follows, buys, bonds, cites, or communicates with.
- Node features: attributes such as age, text embeddings, atom type, or account statistics.
- Edge features: attributes such as relation type, distance, timestamp, or transaction amount.
A GNN learns a vector representation for each node. That vector combines the node’s own features with information from nearby nodes and edges. A readout function then converts representations into the prediction required by the application.
The 2024 Nature Reviews Methods Primers describes GNNs as mathematical models that learn functions over graphs and are a leading approach for predictive models on graph-structured data. The 2021 methods review describes message passing between nodes as the mechanism that captures graph dependence.
How message passing works
Each message-passing layer performs four conceptual operations:
- Compute a message from each neighbor, optionally using the connecting edge’s features.
- Aggregate all incoming messages with a permutation-invariant operation such as sum, mean, or max.
- Combine the aggregate with the node’s current state.
- Apply a learned transformation and nonlinearity to produce the next state.
A generic layer can be written as:
m_v = AGGREGATE({ MESSAGE(h_v, h_u, e_uv) : u in N(v) })
h_v' = UPDATE(h_v, m_v)
After one layer, a node can use one-hop information. After two layers, it can use information from nodes up to two hops away. In practice, adding layers indefinitely is not always helpful: representations can become too similar (over-smoothing), and many distant signals can be compressed into a small number of local vectors (over-squashing).
What can a GNN predict?
| Task | Output | Typical examples | Readout |
|---|---|---|---|
| Node prediction | Label or value per node | User risk, paper topic, atom property | Use the target node embeddings |
| Link prediction | Probability or score for an edge | Friend recommendation, missing knowledge-graph relation | Combine two endpoint embeddings |
| Edge prediction | Label or value per existing edge | Transaction fraud, bond type | Use both endpoints and edge features |
| Graph prediction | One label or value per graph | Molecule toxicity, scene class, fraud subgraph | Pool node embeddings into one graph vector |
Choose the target before choosing the architecture. A model can have excellent node accuracy while being unsuitable for graph-level prediction if its pooling function discards the relevant structure.
GCN, GraphSAGE, GAT, and relational GCN
| Architecture | Main idea | Good starting point when | Watch for |
|---|---|---|---|
| GCN | Normalized neighbor aggregation. | The graph is relatively simple and neighboring nodes tend to have related labels (homophily). | It can blur distinctions and usually assumes one edge type. |
| GraphSAGE | Sample a bounded set of neighbors and aggregate their features. | You need inductive predictions for unseen nodes or graphs, or the graph is too large for full neighborhoods. | Sampling variance and the choice of sampler affect quality. |
| GAT | Learn attention weights that give different neighbors different influence. | Some neighbors are more informative than others. | Attention adds compute and tuning cost; an attention weight is not automatically a causal explanation. |
| Relational GCN | Use relation-specific transformations for typed edges. | Knowledge graphs or other heterogeneous relations are central to the task. | Many relation types increase parameters and can create sparse-data problems. |
Official DGL tutorials cover GCN, GAT, GraphSAGE, and relational GCN implementations. Select by deployment requirements as well as validation score: inductive versus transductive use, homogeneous versus typed edges, graph size, long-range dependencies, calibration, and robustness.
A complete PyTorch Geometric node-classification example
The following script trains a two-layer GCN on the Cora citation dataset. Install dependencies first:
python -m pip install torch torch-geometric
import torch
import torch.nn.functional as F
from torch_geometric.datasets import Planetoid
from torch_geometric.nn import GCNConv
# Cora is a citation graph with node features, edges, labels, and masks.
dataset = Planetoid(root="data/Planetoid", name="Cora")
data = dataset[0]
class GCN(torch.nn.Module):
def __init__(self, in_channels, hidden_channels, out_channels):
super().__init__()
self.conv1 = GCNConv(in_channels, hidden_channels)
self.conv2 = GCNConv(hidden_channels, out_channels)
def forward(self, x, edge_index):
x = self.conv1(x, edge_index)
x = F.relu(x)
x = F.dropout(x, p=0.5, training=self.training)
return self.conv2(x, edge_index)
model = GCN(dataset.num_features, 64, dataset.num_classes)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
for epoch in range(1, 201):
model.train()
optimizer.zero_grad()
logits = model(data.x, data.edge_index)
loss = F.cross_entropy(logits[data.train_mask], data.y[data.train_mask])
loss.backward()
optimizer.step()
if epoch % 20 == 0:
model.eval()
pred = logits.argmax(dim=-1)
test_acc = (pred[data.test_mask] == data.y[data.test_mask]).float().mean()
print(f"epoch={epoch:03d} loss={loss.item():.4f} test_acc={test_acc.item():.4f}")
For a real project, replace the random or benchmark split with a split that matches deployment. A temporal split is usually safer for evolving interactions; a graph-level split is needed when entire graphs, rather than nodes in one graph, are the independent examples.
Building a GNN workflow
- Define the graph. Decide what counts as a node and edge, whether edges are directed, and which relationships are trustworthy.
- Define the target. State whether the output is node-, edge-, or graph-level and identify the prediction time.
- Prevent leakage. Remove features computed after the prediction time. For link prediction, split positive and negative edges carefully so validation does not reveal test edges through topology.
- Build a non-graph baseline. Compare against logistic regression, gradient-boosted trees, or an MLP using only local features.
- Choose features and relations. Normalize numeric features, encode categories, and preserve edge types when they carry meaning.
- Choose an architecture. Start with GCN for a simple homogeneous graph, GraphSAGE for inductive or sampled training, GAT for unequal neighbor importance, and relational GCN for typed edges.
- Choose the evaluation metric. Use macro-F1 or per-class recall for imbalanced node classes, ranking metrics for recommendations, and calibration measurements when scores drive decisions.
- Stress test the structure. Remove or perturb edges, mask features, test new nodes, and evaluate on later time periods or new graph populations.
Can GNNs handle large graphs?
Yes, but usually through sampling, partitioning, or mini-batching rather than loading every neighborhood into every layer. PyTorch Geometric provides loaders for many small graphs and single large graphs, along with multi-GPU and compilation support. DGL documents sparse kernels, auto-batching, CPU and multi-GPU training, and workflows intended for graphs with hundreds of millions of nodes and edges. Those are framework capabilities; actual throughput depends on feature sizes, degree distribution, hardware, and sampling strategy.
Scaling techniques
- Neighbor sampling: sample a fixed number of neighbors per layer, as in GraphSAGE.
- Layer-wise sampling: select a bounded set of nodes for each layer instead of expanding every neighborhood.
- Cluster or partition training: divide a large graph into subgraphs while retaining boundary information.
- Subgraph or temporal batches: train on local regions or time windows when the application permits it.
- Sparse representations: keep adjacency and features in sparse formats where supported.
Sampling reduces memory but changes the learning problem. Monitor the degree distribution of sampled neighbors, increase sample sizes for high-variance regions, and compare sampled validation against full-neighborhood evaluation when possible.
Limitations and failure modes
- Over-smoothing: deep layers can make node embeddings indistinguishable. Use residual or skip connections, normalization, fewer layers, or architectures designed for deeper propagation.
- Over-squashing: exponentially many distant signals are compressed into fixed-size vectors. Reconsider the graph, add useful shortcuts, use hierarchical pooling, or evaluate a global-context model.
- Limited structural expressiveness: standard message-passing models have bounds related to Weisfeiler–Lehman graph tests and may fail to distinguish some structurally different graphs.
- Heterophily: neighboring nodes may intentionally have different labels. Plain neighbor averaging can then hurt; use relation-aware, signed, or heterophily-oriented designs and compare with non-graph baselines.
- Changing graphs: new nodes, deleted edges, and delayed labels require time-aware features and evaluation.
- Bad topology: missing, biased, or adversarial edges can materially change predictions. Report sensitivity to edge perturbations.
- Calibration: high classification accuracy does not guarantee reliable probabilities. Measure calibration and uncertainty if predictions trigger actions.
Debugging checklist
| Symptom | Likely cause | Fix |
|---|---|---|
| Training accuracy rises but validation collapses | Leakage, overfitting, or a split that does not match deployment | Rebuild masks by time or graph, remove future features, add regularization, and compare with a simple baseline. |
| Loss becomes NaN | Unscaled features, exploding gradients, or an invalid edge/feature value | Normalize inputs, lower the learning rate, clip gradients, and check tensors for NaN or infinity. |
| GPU memory is exhausted | Full-neighborhood expansion or high-dimensional features | Use neighbor sampling, smaller batches, fewer layers, mixed precision where appropriate, or graph partitioning. |
| Predictions are almost identical | Over-smoothing, excessive depth, or a disconnected feature pipeline | Inspect embedding variance, reduce depth, add residual paths, and verify that edge_index and features align. |
| New nodes cannot be scored | Transductive training or unavailable features/neighbors | Use an inductive architecture such as GraphSAGE and define an inference-time neighborhood policy. |
| Results change sharply after small edge edits | Topology sensitivity or noisy relationships | Run edge-perturbation tests, add robust training, and treat uncertainty as a deployment signal. |
Performance, reliability, and cost considerations
- Compute: message passing cost grows with the number of sampled or processed edges, feature width, and layer count. Profile data loading and sampling as well as matrix operations.
- Memory: activations for every sampled node and edge can dominate training. Reduce fanout, use checkpointing, or partition the graph.
- Latency: online inference should define a bounded neighborhood and cache stable features. Recomputing a large multi-hop neighborhood per request is often impractical.
- Reproducibility: record random seeds, graph snapshots, feature-generation code, split definitions, sampler settings, and library versions.
- Reliability: monitor graph freshness, missing edges, degree shifts, feature drift, class imbalance, and confidence calibration after deployment.
- Cost: benchmark end-to-end training and inference on representative graph sizes. A model with fewer parameters can still be expensive if it expands many neighbors.
Or skip the browser setup
Graph work often includes visual checks of data portals, experiment dashboards, or model documentation. If you need a clean website image in an automated pipeline, ScreenshotNeo provides a single HTTP request instead of maintaining a browser.
See the ScreenshotNeo API documentation for all options. This example returns a WebP screenshot:
curl -G "https://api.screenshotneo.com/v1/shot" -d access_key=YOUR_API_KEY --data-urlencode url=https://stripe.com -o shot.webp
import requests
r = requests.get("https://api.screenshotneo.com/v1/shot", params={"access_key": "YOUR_API_KEY", "url": "https://stripe.com"}, timeout=90)
open("shot.webp", "wb").write(r.content)
const q = new URLSearchParams({ access_key: 'YOUR_API_KEY', url: 'https://stripe.com' });
const res = await fetch(`https://api.screenshotneo.com/v1/shot?${q}`);
Cookie banners, popups, and chat widgets are removed before the shot. Bot checks, blank pages, and failed loads are never billed, and response headers identify the page verdict and billing status. An MCP server lets AI agents use take_screenshot, get_page_info, and capture_pdf. The Free plan includes 1,000 screenshots a month with no card; paid plans start at $5 for 3,000 shots. Create a free ScreenshotNeo account.
Frequently asked questions
Do GNNs require one connected graph?
No. You can train on many separate graphs for graph-level tasks, or on one large graph for node and link tasks. The batching strategy changes, but the message-passing idea is the same.
Are attention weights an explanation?
They show how a particular GAT layer weighted neighbors, but they are not a guaranteed causal explanation. Validate explanations with perturbation tests or a separate explanation method.
When should I avoid a GNN?
A GNN may be a poor fit when relationships are unreliable, the target has no graph dependence, the graph changes too quickly to maintain, or a tabular baseline already meets the requirement.
Can a GNN use text or images?
Yes. Encode text, images, or other modalities into node or edge features, then let message passing combine those representations with graph structure.
What should I read next?
William L. Hamilton’s Graph Representation Learning (Springer, 2020) covers GNN models, practice, and theoretical motivations, with applications including chemical synthesis, 3D vision, recommender systems, question answering, and social-network analysis.


