PyTorch Geometric (PyG) — Graph Neural Networks
Overview
PyTorch Geometric is a library built on PyTorch for developing and training Graph Neural Networks (GNNs). It provides 40+ convolutional layers, mini-batch processing via block-diagonal adjacency matrices, neighbor sampling for large-scale graphs, and heterogeneous graph support for multi-type node/edge networks.
When to Use
- Node classification on citation, social, or biological networks
- Graph-level classification (molecular activity, protein function)
- Link prediction (knowledge graphs, recommendation systems)
- Molecular property prediction (drug discovery, quantum chemistry)
- 3D point cloud processing and mesh analysis
- Large-scale graph learning with neighbor sampling (>100K nodes)
- Heterogeneous graphs with multiple node/edge types
- For non-graph deep learning → use PyTorch directly
- For traditional graph algorithms (shortest path, centrality) → use NetworkX
Prerequisites
pip install torch torch_geometric
# Optional sparse operations (recommended):
# pip install pyg_lib torch_scatter torch_sparse torch_cluster
import torch
import torch.nn.functional as F
from torch_geometric.data import Data
from torch_geometric.nn import GCNConv
Quick Start
from torch_geometric.datasets import Planetoid
from torch_geometric.nn import GCNConv
import torch, torch.nn.functional as F
dataset = Planetoid(root='/tmp/Cora', name='Cora')
data = dataset[0]
class GCN(torch.nn.Module):
def __init__(self):
super().__init__()
self.conv1 = GCNConv(dataset.num_features, 16)
self.conv2 = GCNConv(16, dataset.num_classes)
def forward(self, data):
x = F.relu(self.conv1(data.x, data.edge_index))
return self.conv2(x, data.edge_index)
model = GCN()
optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
for epoch in range(200):
model.train(); optimizer.zero_grad()
F.cross_entropy(model(data)[data.train_mask], data.y[data.train_mask]).backward()
optimizer.step()
model.eval()
pred = model(data).argmax(dim=1)
acc = (pred[data.test_mask] == data.y[data.test_mask]).float().mean()
print(f'Test Accuracy: {acc:.4f}') # ~0.81
Core API
1. Data Representation
import torch
from torch_geometric.data import Data
# Create a graph: 3 nodes, 4 edges (undirected)
edge_index = torch.tensor([[0, 1, 1, 2],
[1, 0, 2, 1]], dtype=torch.long)
x = torch.randn(3, 16) # Node features [num_nodes, features]
y = torch.tensor([0, 1, 0]) # Node labels
data = Data(x=x, edge_index=edge_index, y=y)
print(f'Nodes: {data.num_nodes}, Edges: {data.num_edges}')
print(f'Features: {data.num_node_features}')
print(f'Has self-loops: {data.has_self_loops()}')
print(f'Is undirected: {data.is_undirected()}')
# Optional attributes
data.edge_attr = torch.randn(4, 8) # Edge features [num_edges, features]
data.pos = torch.randn(3, 3) # Node positions (3D)
data.train_mask = torch.tensor([True, True, False]) # Custom masks
# Mini-batch processing — graphs concatenated as block-diagonal
from torch_geometric.loader import DataLoader
loader = DataLoader(dataset, batch_size=32, shuffle=True)
for batch in loader:
print(f'Graphs: {batch.num_graphs}, Nodes: {batch.num_nodes}')
# batch.batch maps each node → its source graph index
# No padding needed — computationally efficient
2. Convolutional Layers
from torch_geometric.nn import GCNConv, GATConv, SAGEConv, GINConv
import torch.nn as nn
# GCNConv — spectral graph convolution (baseline)
conv = GCNConv(in_channels=16, out_channels=32)
# Supports: edge_weight, SparseTensor, Bipartite, Lazy init
# GATConv — attention-based neighbor weighting
conv = GATConv(16, 32, heads=8, dropout=0.6)
# Output: [N, heads * out_channels] (concat) or [N, out_channels] (concat=False)
# SAGEConv — inductive learning via sampling
conv = SAGEConv(16, 32, aggr='mean') # 'mean', 'max', 'lstm'
# GINConv — maximally powerful for graph isomorphism
nn_module = nn.Sequential(nn.Linear(16, 32), nn.ReLU(), nn.Linear(32, 32))
conv = GINConv(nn_module)
# TransformerConv — graph transformer
from torch_geometric.nn import TransformerConv
conv = TransformerConv(16, 32, heads=8, beta=True)
# All layers: x_out = conv(x, edge_index)
x_out = conv(x, edge_index)
print(f'Output shape: {x_out.shape}') # [num_nodes, out_channels]
3. Custom Message Passing
from torch_geometric.nn import MessagePassing
from torch_geometric.utils import add_self_loops, degree
class CustomConv(MessagePassing):
def __init__(self, in_channels, out_channels):
super().__init__(aggr='add') # 'add', 'mean', 'max'
self.lin = torch.nn.Linear(in_channels, out_channels)
def forward(self, x, edge_index):
edge_index, _ = add_self_loops(edge_index, num_nodes=x.size(0))
x = self.lin(x)
# Degree-based normalization
row, col = edge_index
deg = degree(col, x.size(0), dtype=x.dtype)
norm = deg.pow(-0.5)
norm = norm[row] * norm[col]
return self.propagate(edge_index, x=x, norm=norm)
def message(self, x_j, norm):
# x_j: source node features (automatic via _j suffix)
return norm.view(-1, 1) * x_j
# Key methods: forward(), message(), aggregate(), update()
# _i suffix → target node, _j suffix → source node
4. Pooling & Graph-Level Readout
from torch_geometric.nn import (
global_mean_pool, global_max_pool, global_add_pool,
TopKPooling, SAGPooling
)
# Global pooling: node features → graph-level representation
x_graph = global_mean_pool(x, batch) # [num_graphs, features]
# Hierarchical pooling: coarsen graph
pool = TopKPooling(64, ratio=0.8) # Keep top 80% nodes
x, edge_index, _, batch, _, _ = pool(x, edge_index, None, batch)
# Graph classification model
class GraphClassifier(torch.nn.Module):
def __init__(self, num_features, num_classes):
super().__init__()
self.conv1 = GCNConv(num_features, 64)
self.conv2 = GCNConv(64, 64)
self.pool = TopKPooling(64, ratio=0.8)
self.lin = torch.nn.Linear(64, num_classes)
def forward(self, data):
x, edge_index, batch = data.x, data.edge_index, data.batch
x = F.relu(self.conv1(x, edge_index))
x, edge_index, _, batch, _, _ = self.pool(x, edge_index, None, batch)
x = F.relu(self.conv2(x, edge_index))
x = global_mean_pool(x, batch)
return self.lin(x)
5. Heterogeneous Graphs
from torch_geometric.data import HeteroData
from torch_geometric.nn import HeteroConv, GCNConv, SAGEConv, to_hetero
# Create heterogeneous graph
data = HeteroData()
data['paper'].x = torch.randn(100, 128)
data['author'].x = torch.randn(200, 64)
data['author', 'writes', 'paper'].edge_index = torch.randint(0, 200, (2, 500))
data['paper', 'cites', 'paper'].edge_index = torch.randint(0, 100, (2, 300))
print(data) # Shows all node/edge types
# Method 1: Auto-convert homogeneous model
model = GCN(...)
model = to_hetero(model, data.metadata(), aggr='sum')
out = model(data.x_dict, data.edge_index_dict)
# Method 2: Custom per-edge-type convolutions
class HeteroGNN(torch.nn.Module):
def __init__(self):
super().__init__()
self.conv1 = HeteroConv({
('paper', 'cites', 'paper'): GCNConv(-1, 64),
('author', 'writes', 'paper'): SAGEConv((-1, -1), 64),
}, aggr='sum')
def forward(self, x_dict, edge_index_dict):
x_dict = self.conv1(x_dict, edge_index_dict)
return {k: F.relu(v) for k, v in x_dict.items()}
6. Transforms & Preprocessing
from torch_geometric.transforms import (
NormalizeFeatures, AddSelfLoops, ToUndirected,
RandomNodeSplit, RandomLinkSplit, Compose,
KNNGraph, RadiusGraph, AddLaplacianEigenvectorPE
)
# Single transform
dataset = Planetoid(root='/tmp/Cora', name='Cora', transform=NormalizeFeatures())
# Compose multiple transforms
transform = Compose([
ToUndirected(),
AddSelfLoops(),
NormalizeFeatures(),
])
# Data splitting
node_split = RandomNodeSplit(num_val=0.1, num_test=0.2)
link_split = RandomLinkSplit(num_val=0.1, num_test=0.2, is_undirected=True)
# Point cloud → graph
pc_transform = Compose([KNNGraph(k=6), NormalizeFeatures()])
# Positional encodings (for Graph Transformers)
pe_transform = AddLaplacianEigenvectorPE(k=10)
Key Concepts
Layer Selection Guide
| Task |
Layer |
Key Feature |
| Baseline / general |
GCNConv |
Spectral, cached, edge_weight |
| Variable neighbor importance |
GATConv / GATv2Conv |
Multi-head attention |
| Large-scale inductive |
SAGEConv |
Sampling-friendly, mean/max/lstm aggr |
| Graph classification |
GINConv |
Maximally powerful WL-test |
| Long-range dependencies |
TransformerConv |
Graph transformer |
| Spectral filtering |
ChebConv |
Chebyshev polynomials, K hops |
| Rich edge features |
NNConv |
Edge NN processes edge_attr |
| Molecular / 3D structures |
SchNet, DimeNet |
Continuous filters, angles |
| Heterogeneous / multi-relation |
RGCNConv, HGTConv |
Multiple edge types |
| Point clouds |
EdgeConv, PointNetConv |
Dynamic graphs, local features |
| Deep GNNs (avoid oversmoothing) |
APPNP + PairNorm |
Separated propagation |
Data Flow Architecture
- edge_index:
[2, num_edges] COO format. Row 0 = source, Row 1 = target
- Mini-batch: Block-diagonal adjacency +
batch vector mapping nodes → graphs. No padding
- Neighbor sampling:
NeighborLoader samples K-hop subgraphs per seed node. Output is directed, relabeled
- Heterogeneous:
x_dict (per-type features), edge_index_dict (per-relation edges), metadata() for schema
Aggregation Options
| Aggregation |
Class |
Use Case |
| Sum |
SumAggregation |
Counting-sensitive tasks |
| Mean |
MeanAggregation |
Degree-invariant |
| Max |
MaxAggregation |
Salient feature detection |
| Softmax |
SoftmaxAggregation(learn=True) |
Learnable attention |
| Multi |
MultiAggregation(['mean','max','std']) |
Combined signals |
Common Workflows
1. Node Classification (Full Graph)
import torch
import torch.nn.functional as F
from torch_geometric.datasets import Planetoid
from torch_geometric.nn import GCNConv
dataset = Planetoid(root='/tmp/Cora', name='Cora')
data = dataset[0]
class GCN(torch.nn.Module):
def __init__(self):
super().__init__()
self.conv1 = GCNConv(dataset.num_features, 16)
self.conv2 = GCNConv(16, dataset.num_classes)
def forward(self, data):
x = F.dropout(F.relu(self.conv1(data.x, data.edge_index)), p=0.5, training=self.training)
return self.conv2(x, data.edge_index)
model = GCN()
optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
# Training
for epoch in range(200):
model.train(); optimizer.zero_grad()
out = model(data)
F.cross_entropy(out[data.train_mask], data.y[data.train_mask]).backward()
optimizer.step()
# Evaluation
model.eval()
pred = model(data).argmax(dim=1)
acc = (pred[data.test_mask] == data.y[data.test_mask]).float().mean()
print(f'Test Accuracy: {acc:.4f}')
2. Graph Classification (Mini-Batch)
from torch_geometric.datasets import TUDataset
from torch_geometric.loader import DataLoader
from torch_geometric.nn import GCNConv, global_mean_pool
dataset = TUDataset(root='/tmp/ENZYMES', name='ENZYMES')
train_dataset = dataset[:int(0.8 * len(dataset))]
test_dataset = dataset[int(0.8 * len(dataset)):]
train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
class GraphNet(torch.nn.Module):
def __init__(self):
super().__init__()
self.conv1 = GCNConv(dataset.num_features, 64)
self.conv2 = GCNConv(64, 64)
self.lin = torch.nn.Linear(64, dataset.num_classes)
def forward(self, data):
x = F.relu(self.conv1(data.x, data.edge_index))
x = F.relu(self.conv2(x, data.edge_index))
x = global_mean_pool(x, data.batch)
return self.lin(x)
model = GraphNet()
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
for epoch in range(100):
model.train()
for batch in train_loader:
optimizer.zero_grad()
F.cross_entropy(model(batch), batch.y).backward()
optimizer.step()
3. Large-Scale with Neighbor Sampling
from torch_geometric.loader import NeighborLoader
# Sample 25 1-hop and 10 2-hop neighbors per seed node
train_loader = NeighborLoader(
data,
num_neighbors=[25, 10],
batch_size=128,
input_nodes=data.train_mask,
)
model.train()
for batch in train_loader:
optimizer.zero_grad()
out = model(batch)
# Only compute loss on seed nodes (first batch_size nodes)
loss = F.cross_entropy(out[:batch.batch_size], batch.y[:batch.batch_size])
loss.backward()
optimizer.step()
# Note: output subgraphs are directed, indices relabeled 0..N-1
Key Parameters
| Parameter |
Module |
Default |
Range |
Effect |
in_channels |
All Conv layers |
— |
int |
Input feature dimension |
out_channels |
All Conv layers |
— |
int |
Output feature dimension |
heads |
GATConv |
1 |
1-16 |
Number of attention heads |
dropout |
GATConv |
0.0 |
0-0.8 |
Attention weight dropout |
aggr |
MessagePassing |
'add' |
add/mean/max |
Neighbor aggregation |
K |
ChebConv |
— |
2-5 |
Chebyshev polynomial order |
num_neighbors |
NeighborLoader |
— |
list[int] |
Neighbors per hop (e.g., [25,10]) |
batch_size |
DataLoader |
— |
16-512 |
Graphs or seed nodes per batch |
ratio |
TopKPooling |
0.5 |
0.1-0.9 |
Fraction of nodes to keep |
lr |
Adam |
— |
1e-4 to 0.01 |
Learning rate |
weight_decay |
Adam |
0 |
0 to 5e-3 |
L2 regularization |
Best Practices
- Start with GCNConv: Use 2-layer GCN as baseline before trying complex architectures
- Use lazy initialization: Pass
-1 as in_channels to infer dimensions automatically: GCNConv(-1, 64)
- Normalize features: Apply
NormalizeFeatures() transform for citation/social networks
- Anti-pattern — too many layers: GNNs typically need only 2-3 layers. Deeper causes oversmoothing. Use
JumpingKnowledge or PairNorm if you need depth
- GPU transfer: Move both model AND data to GPU:
model.to(device), data.to(device)
- Anti-pattern — ignoring batch vector: In graph classification, always use
global_mean_pool(x, batch) — forgetting batch pools across all graphs
Common Recipes
Recipe: Model Explainability (GNNExplainer)
from torch_geometric.explain import Explainer, GNNExplainer
explainer = Explainer(
model=model,
algorithm=GNNExplainer(epochs=200),
explanation_type='model',
node_mask_type='attributes',
edge_mask_type='object',
model_config=dict(mode='multiclass_classification', task_level='node', return_type='log_probs'),
)
explanation = explainer(data.x, data.edge_index, index=10)
print(f'Important edges: {explanation.edge_mask.topk(5).indices}')
print(f'Important features: {explanation.node_mask[10].topk(5).indices}')
Recipe: Custom InMemoryDataset
from torch_geometric.data import InMemoryDataset, Data
class MyDataset(InMemoryDataset):
def __init__(self, root, transform=None, pre_transform=None):
super().__init__(root, transform, pre_transform)
self.load(self.processed_paths[0])
@property
def raw_file_names(self):
return ['data.csv']
@property
def processed_file_names(self):
return ['data.pt']
def process(self):
data_list = []
# Build Data objects from raw files
edge_index = torch.tensor([[0, 1], [1, 0]], dtype=torch.long)
x = torch.randn(2, 16)
data_list.append(Data(x=x, edge_index=edge_index, y=torch.tensor([0])))
if self.pre_filter is not None:
data_list = [d for d in data_list if self.pre_filter(d)]
if self.pre_transform is not None:
data_list = [self.pre_transform(d) for d in data_list]
self.save(data_list, self.processed_paths[0])
Recipe: Deep GNN with JumpingKnowledge
from torch_geometric.nn import GCNConv, JumpingKnowledge, LayerNorm
class DeepGNN(torch.nn.Module):
def __init__(self, in_ch, hidden, num_layers, out_ch):
super().__init__()
self.convs = torch.nn.ModuleList()
self.norms = torch.nn.ModuleList()
self.convs.append(GCNConv(in_ch, hidden))
self.norms.append(LayerNorm(hidden))
for _ in range(num_layers - 2):
self.convs.append(GCNConv(hidden, hidden))
self.norms.append(LayerNorm(hidden))
self.convs.append(GCNConv(hidden, hidden))
self.jk = JumpingKnowledge(mode='cat')
self.lin = torch.nn.Linear(hidden * num_layers, out_ch)
def forward(self, x, edge_index, batch):
xs = []
for conv, norm in zip(self.convs[:-1], self.norms):
x = F.relu(norm(conv(x, edge_index)))
xs.append(x)
xs.append(self.convs[-1](x, edge_index))
return self.lin(global_mean_pool(self.jk(xs), batch))
Troubleshooting
| Problem |
Cause |
Solution |
edge_index shape error |
Wrong format (should be [2, E]) |
Ensure COO format: torch.tensor([[src...],[dst...]], dtype=torch.long) |
| OOM on large graph |
Full-graph forward pass |
Use NeighborLoader for mini-batch training |
| Low accuracy |
Oversmoothing (too many layers) |
Reduce to 2-3 layers, add JumpingKnowledge or PairNorm |
| NaN in training |
Exploding gradients |
Add gradient clipping, reduce learning rate, check feature scale |
| Wrong graph-level output |
Missing batch in pooling |
Pass batch tensor to global_mean_pool(x, batch) |
| Heterogeneous type error |
Mismatched node/edge types |
Check data.metadata() matches model definition |
| Slow DataLoader |
Large graph, no sampling |
Use NeighborLoader with reasonable num_neighbors (e.g., [25,10]) |
x dimension mismatch |
Multi-head attention output |
For GATConv: output is heads*out_channels unless concat=False |
| Import error for sparse ops |
Missing optional dependencies |
Install torch_scatter, torch_sparse from PyG wheels |
| Pre-transform not applied |
Dataset already processed |
Delete processed/ directory and reload |
Bundled Resources
references/layers_transforms_reference.md — Complete catalog of 40+ convolutional layers (GCN, GAT, SAGE, GIN, molecular layers, hypergraph), aggregation operators, pooling (global + hierarchical), normalization layers, pre-built models, auto-encoders, knowledge graph embeddings, utility layers. Transform catalog: structure, feature, spatial, augmentation, mesh, specialized. Consolidated from original layers_reference.md (486 lines) + transforms_reference.md (680 lines). Script functionality (benchmark_model.py, create_gnn_template.py, visualize_graph.py) covered by Core API code blocks and Common Recipes
references/datasets_catalog.md — Comprehensive dataset catalog organized by domain: citation networks (Planetoid, Coauthor, Amazon), graph classification (TUDataset 120+ benchmarks), molecular (QM9, ZINC, MoleculeNet), social (Reddit, Twitch), knowledge graphs (WordNet, FB15k), heterogeneous (OGB_MAG, MovieLens, DBLP), temporal (JODIE), 3D meshes (ShapeNet, ModelNet), OGB integration. Consolidated from original datasets_reference.md (575 lines)
Related Skills
- matplotlib-scientific-plotting — Visualize graph structures, training curves, attention weights
References
1---2name: torch-geometric-graph-neural-networks3description: PyTorch Geometric (PyG) for graph neural networks: node/graph classification, link prediction with GCN, GAT, GraphSAGE, GIN. Message passing, mini-batches, heterogeneous graphs, neighbor sampling, explainability. Supports molecules (QM9, MoleculeNet), social/knowledge graphs, 3D point clouds. For non-graph DL use PyTorch; for classical graph algorithms use NetworkX.4license: MIT5---6
7# PyTorch Geometric (PyG) — Graph Neural Networks
8
9## Overview
10
11PyTorch Geometric is a library built on PyTorch for developing and training Graph Neural Networks (GNNs). It provides 40+ convolutional layers, mini-batch processing via block-diagonal adjacency matrices, neighbor sampling for large-scale graphs, and heterogeneous graph support for multi-type node/edge networks.
12
13## When to Use
14
15- Node classification on citation, social, or biological networks
16- Graph-level classification (molecular activity, protein function)
17- Link prediction (knowledge graphs, recommendation systems)
18- Molecular property prediction (drug discovery, quantum chemistry)
19- 3D point cloud processing and mesh analysis
20- Large-scale graph learning with neighbor sampling (>100K nodes)
21- Heterogeneous graphs with multiple node/edge types
22- **For non-graph deep learning** → use PyTorch directly
23- **For traditional graph algorithms (shortest path, centrality)** → use NetworkX
24
25## Prerequisites
26
27```bash
28pip install torch torch_geometric
29# Optional sparse operations (recommended):
30# pip install pyg_lib torch_scatter torch_sparse torch_cluster
31```
32
33```python
34import torch
35import torch.nn.functional as F
36from torch_geometric.data import Data
37from torch_geometric.nn import GCNConv
38```
39
40## Quick Start
41
42```python
43from torch_geometric.datasets import Planetoid
44from torch_geometric.nn import GCNConv
45import torch, torch.nn.functional as F
46
47dataset = Planetoid(root='/tmp/Cora', name='Cora')
48data = dataset[0]
49
50class GCN(torch.nn.Module):
51 def __init__(self):
52 super().__init__()
53 self.conv1 = GCNConv(dataset.num_features, 16)
54 self.conv2 = GCNConv(16, dataset.num_classes)
55 def forward(self, data):
56 x = F.relu(self.conv1(data.x, data.edge_index))
57 return self.conv2(x, data.edge_index)
58
59model = GCN()
60optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
61for epoch in range(200):
62 model.train(); optimizer.zero_grad()
63 F.cross_entropy(model(data)[data.train_mask], data.y[data.train_mask]).backward()
64 optimizer.step()
65
66model.eval()
67pred = model(data).argmax(dim=1)
68acc = (pred[data.test_mask] == data.y[data.test_mask]).float().mean()
69print(f'Test Accuracy: {acc:.4f}') # ~0.81
70```
71
72## Core API
73
74### 1. Data Representation
75
76```python
77import torch
78from torch_geometric.data import Data
79
80# Create a graph: 3 nodes, 4 edges (undirected)
81edge_index = torch.tensor([[0, 1, 1, 2],
82 [1, 0, 2, 1]], dtype=torch.long)
83x = torch.randn(3, 16) # Node features [num_nodes, features]
84y = torch.tensor([0, 1, 0]) # Node labels
85
86data = Data(x=x, edge_index=edge_index, y=y)
87print(f'Nodes: {data.num_nodes}, Edges: {data.num_edges}')
88print(f'Features: {data.num_node_features}')
89print(f'Has self-loops: {data.has_self_loops()}')
90print(f'Is undirected: {data.is_undirected()}')
91
92# Optional attributes
93data.edge_attr = torch.randn(4, 8) # Edge features [num_edges, features]
94data.pos = torch.randn(3, 3) # Node positions (3D)
95data.train_mask = torch.tensor([True, True, False]) # Custom masks
96```
97
98```python
99# Mini-batch processing — graphs concatenated as block-diagonal
100from torch_geometric.loader import DataLoader
101
102loader = DataLoader(dataset, batch_size=32, shuffle=True)
103for batch in loader:
104 print(f'Graphs: {batch.num_graphs}, Nodes: {batch.num_nodes}')
105 # batch.batch maps each node → its source graph index
106 # No padding needed — computationally efficient
107```
108
109### 2. Convolutional Layers
110
111```python
112from torch_geometric.nn import GCNConv, GATConv, SAGEConv, GINConv
113import torch.nn as nn
114
115# GCNConv — spectral graph convolution (baseline)
116conv = GCNConv(in_channels=16, out_channels=32)
117# Supports: edge_weight, SparseTensor, Bipartite, Lazy init
118
119# GATConv — attention-based neighbor weighting
120conv = GATConv(16, 32, heads=8, dropout=0.6)
121# Output: [N, heads * out_channels] (concat) or [N, out_channels] (concat=False)
122
123# SAGEConv — inductive learning via sampling
124conv = SAGEConv(16, 32, aggr='mean') # 'mean', 'max', 'lstm'
125
126# GINConv — maximally powerful for graph isomorphism
127nn_module = nn.Sequential(nn.Linear(16, 32), nn.ReLU(), nn.Linear(32, 32))
128conv = GINConv(nn_module)
129
130# TransformerConv — graph transformer
131from torch_geometric.nn import TransformerConv
132conv = TransformerConv(16, 32, heads=8, beta=True)
133
134# All layers: x_out = conv(x, edge_index)
135x_out = conv(x, edge_index)
136print(f'Output shape: {x_out.shape}') # [num_nodes, out_channels]
137```
138
139### 3. Custom Message Passing
140
141```python
142from torch_geometric.nn import MessagePassing
143from torch_geometric.utils import add_self_loops, degree
144
145class CustomConv(MessagePassing):
146 def __init__(self, in_channels, out_channels):
147 super().__init__(aggr='add') # 'add', 'mean', 'max'
148 self.lin = torch.nn.Linear(in_channels, out_channels)
149
150 def forward(self, x, edge_index):
151 edge_index, _ = add_self_loops(edge_index, num_nodes=x.size(0))
152 x = self.lin(x)
153
154 # Degree-based normalization
155 row, col = edge_index
156 deg = degree(col, x.size(0), dtype=x.dtype)
157 norm = deg.pow(-0.5)
158 norm = norm[row] * norm[col]
159
160 return self.propagate(edge_index, x=x, norm=norm)
161
162 def message(self, x_j, norm):
163 # x_j: source node features (automatic via _j suffix)
164 return norm.view(-1, 1) * x_j
165
166# Key methods: forward(), message(), aggregate(), update()
167# _i suffix → target node, _j suffix → source node
168```
169
170### 4. Pooling & Graph-Level Readout
171
172```python
173from torch_geometric.nn import (
174 global_mean_pool, global_max_pool, global_add_pool,
175 TopKPooling, SAGPooling
176)
177
178# Global pooling: node features → graph-level representation
179x_graph = global_mean_pool(x, batch) # [num_graphs, features]
180
181# Hierarchical pooling: coarsen graph
182pool = TopKPooling(64, ratio=0.8) # Keep top 80% nodes
183x, edge_index, _, batch, _, _ = pool(x, edge_index, None, batch)
184
185# Graph classification model
186class GraphClassifier(torch.nn.Module):
187 def __init__(self, num_features, num_classes):
188 super().__init__()
189 self.conv1 = GCNConv(num_features, 64)
190 self.conv2 = GCNConv(64, 64)
191 self.pool = TopKPooling(64, ratio=0.8)
192 self.lin = torch.nn.Linear(64, num_classes)
193
194 def forward(self, data):
195 x, edge_index, batch = data.x, data.edge_index, data.batch
196 x = F.relu(self.conv1(x, edge_index))
197 x, edge_index, _, batch, _, _ = self.pool(x, edge_index, None, batch)
198 x = F.relu(self.conv2(x, edge_index))
199 x = global_mean_pool(x, batch)
200 return self.lin(x)
201```
202
203### 5. Heterogeneous Graphs
204
205```python
206from torch_geometric.data import HeteroData
207from torch_geometric.nn import HeteroConv, GCNConv, SAGEConv, to_hetero
208
209# Create heterogeneous graph
210data = HeteroData()
211data['paper'].x = torch.randn(100, 128)
212data['author'].x = torch.randn(200, 64)
213data['author', 'writes', 'paper'].edge_index = torch.randint(0, 200, (2, 500))
214data['paper', 'cites', 'paper'].edge_index = torch.randint(0, 100, (2, 300))
215print(data) # Shows all node/edge types
216
217# Method 1: Auto-convert homogeneous model
218model = GCN(...)
219model = to_hetero(model, data.metadata(), aggr='sum')
220out = model(data.x_dict, data.edge_index_dict)
221```
222
223```python
224# Method 2: Custom per-edge-type convolutions
225class HeteroGNN(torch.nn.Module):
226 def __init__(self):
227 super().__init__()
228 self.conv1 = HeteroConv({
229 ('paper', 'cites', 'paper'): GCNConv(-1, 64),
230 ('author', 'writes', 'paper'): SAGEConv((-1, -1), 64),
231 }, aggr='sum')
232
233 def forward(self, x_dict, edge_index_dict):
234 x_dict = self.conv1(x_dict, edge_index_dict)
235 return {k: F.relu(v) for k, v in x_dict.items()}
236```
237
238### 6. Transforms & Preprocessing
239
240```python
241from torch_geometric.transforms import (
242 NormalizeFeatures, AddSelfLoops, ToUndirected,
243 RandomNodeSplit, RandomLinkSplit, Compose,
244 KNNGraph, RadiusGraph, AddLaplacianEigenvectorPE
245)
246
247# Single transform
248dataset = Planetoid(root='/tmp/Cora', name='Cora', transform=NormalizeFeatures())
249
250# Compose multiple transforms
251transform = Compose([
252 ToUndirected(),
253 AddSelfLoops(),
254 NormalizeFeatures(),
255])
256
257# Data splitting
258node_split = RandomNodeSplit(num_val=0.1, num_test=0.2)
259link_split = RandomLinkSplit(num_val=0.1, num_test=0.2, is_undirected=True)
260
261# Point cloud → graph
262pc_transform = Compose([KNNGraph(k=6), NormalizeFeatures()])
263
264# Positional encodings (for Graph Transformers)
265pe_transform = AddLaplacianEigenvectorPE(k=10)
266```
267
268## Key Concepts
269
270### Layer Selection Guide
271
272| Task | Layer | Key Feature |
273|------|-------|-------------|
274| Baseline / general | `GCNConv` | Spectral, cached, edge_weight |
275| Variable neighbor importance | `GATConv` / `GATv2Conv` | Multi-head attention |
276| Large-scale inductive | `SAGEConv` | Sampling-friendly, mean/max/lstm aggr |
277| Graph classification | `GINConv` | Maximally powerful WL-test |
278| Long-range dependencies | `TransformerConv` | Graph transformer |
279| Spectral filtering | `ChebConv` | Chebyshev polynomials, K hops |
280| Rich edge features | `NNConv` | Edge NN processes edge_attr |
281| Molecular / 3D structures | `SchNet`, `DimeNet` | Continuous filters, angles |
282| Heterogeneous / multi-relation | `RGCNConv`, `HGTConv` | Multiple edge types |
283| Point clouds | `EdgeConv`, `PointNetConv` | Dynamic graphs, local features |
284| Deep GNNs (avoid oversmoothing) | `APPNP` + `PairNorm` | Separated propagation |
285
286### Data Flow Architecture
287
288- **edge_index**: `[2, num_edges]` COO format. Row 0 = source, Row 1 = target
289- **Mini-batch**: Block-diagonal adjacency + `batch` vector mapping nodes → graphs. No padding
290- **Neighbor sampling**: `NeighborLoader` samples K-hop subgraphs per seed node. Output is directed, relabeled
291- **Heterogeneous**: `x_dict` (per-type features), `edge_index_dict` (per-relation edges), `metadata()` for schema
292
293### Aggregation Options
294
295| Aggregation | Class | Use Case |
296|-------------|-------|----------|
297| Sum | `SumAggregation` | Counting-sensitive tasks |
298| Mean | `MeanAggregation` | Degree-invariant |
299| Max | `MaxAggregation` | Salient feature detection |
300| Softmax | `SoftmaxAggregation(learn=True)` | Learnable attention |
301| Multi | `MultiAggregation(['mean','max','std'])` | Combined signals |
302
303## Common Workflows
304
305### 1. Node Classification (Full Graph)
306
307```python
308import torch
309import torch.nn.functional as F
310from torch_geometric.datasets import Planetoid
311from torch_geometric.nn import GCNConv
312
313dataset = Planetoid(root='/tmp/Cora', name='Cora')
314data = dataset[0]
315
316class GCN(torch.nn.Module):
317 def __init__(self):
318 super().__init__()
319 self.conv1 = GCNConv(dataset.num_features, 16)
320 self.conv2 = GCNConv(16, dataset.num_classes)
321 def forward(self, data):
322 x = F.dropout(F.relu(self.conv1(data.x, data.edge_index)), p=0.5, training=self.training)
323 return self.conv2(x, data.edge_index)
324
325model = GCN()
326optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
327
328# Training
329for epoch in range(200):
330 model.train(); optimizer.zero_grad()
331 out = model(data)
332 F.cross_entropy(out[data.train_mask], data.y[data.train_mask]).backward()
333 optimizer.step()
334
335# Evaluation
336model.eval()
337pred = model(data).argmax(dim=1)
338acc = (pred[data.test_mask] == data.y[data.test_mask]).float().mean()
339print(f'Test Accuracy: {acc:.4f}')
340```
341
342### 2. Graph Classification (Mini-Batch)
343
344```python
345from torch_geometric.datasets import TUDataset
346from torch_geometric.loader import DataLoader
347from torch_geometric.nn import GCNConv, global_mean_pool
348
349dataset = TUDataset(root='/tmp/ENZYMES', name='ENZYMES')
350train_dataset = dataset[:int(0.8 * len(dataset))]
351test_dataset = dataset[int(0.8 * len(dataset)):]
352train_loader = DataLoader(train_dataset, batch_size=32, shuffle=True)
353
354class GraphNet(torch.nn.Module):
355 def __init__(self):
356 super().__init__()
357 self.conv1 = GCNConv(dataset.num_features, 64)
358 self.conv2 = GCNConv(64, 64)
359 self.lin = torch.nn.Linear(64, dataset.num_classes)
360 def forward(self, data):
361 x = F.relu(self.conv1(data.x, data.edge_index))
362 x = F.relu(self.conv2(x, data.edge_index))
363 x = global_mean_pool(x, data.batch)
364 return self.lin(x)
365
366model = GraphNet()
367optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
368for epoch in range(100):
369 model.train()
370 for batch in train_loader:
371 optimizer.zero_grad()
372 F.cross_entropy(model(batch), batch.y).backward()
373 optimizer.step()
374```
375
376### 3. Large-Scale with Neighbor Sampling
377
378```python
379from torch_geometric.loader import NeighborLoader
380
381# Sample 25 1-hop and 10 2-hop neighbors per seed node
382train_loader = NeighborLoader(
383 data,
384 num_neighbors=[25, 10],
385 batch_size=128,
386 input_nodes=data.train_mask,
387)
388
389model.train()
390for batch in train_loader:
391 optimizer.zero_grad()
392 out = model(batch)
393 # Only compute loss on seed nodes (first batch_size nodes)
394 loss = F.cross_entropy(out[:batch.batch_size], batch.y[:batch.batch_size])
395 loss.backward()
396 optimizer.step()
397# Note: output subgraphs are directed, indices relabeled 0..N-1
398```
399
400## Key Parameters
401
402| Parameter | Module | Default | Range | Effect |
403|-----------|--------|---------|-------|--------|
404| `in_channels` | All Conv layers | — | int | Input feature dimension |
405| `out_channels` | All Conv layers | — | int | Output feature dimension |
406| `heads` | GATConv | 1 | 1-16 | Number of attention heads |
407| `dropout` | GATConv | 0.0 | 0-0.8 | Attention weight dropout |
408| `aggr` | MessagePassing | 'add' | add/mean/max | Neighbor aggregation |
409| `K` | ChebConv | — | 2-5 | Chebyshev polynomial order |
410| `num_neighbors` | NeighborLoader | — | list[int] | Neighbors per hop (e.g., [25,10]) |
411| `batch_size` | DataLoader | — | 16-512 | Graphs or seed nodes per batch |
412| `ratio` | TopKPooling | 0.5 | 0.1-0.9 | Fraction of nodes to keep |
413| `lr` | Adam | — | 1e-4 to 0.01 | Learning rate |
414| `weight_decay` | Adam | 0 | 0 to 5e-3 | L2 regularization |
415
416## Best Practices
417
4181. **Start with GCNConv**: Use 2-layer GCN as baseline before trying complex architectures
4192. **Use lazy initialization**: Pass `-1` as `in_channels` to infer dimensions automatically: `GCNConv(-1, 64)`
4203. **Normalize features**: Apply `NormalizeFeatures()` transform for citation/social networks
4214. **Anti-pattern — too many layers**: GNNs typically need only 2-3 layers. Deeper causes oversmoothing. Use `JumpingKnowledge` or `PairNorm` if you need depth
4225. **GPU transfer**: Move both model AND data to GPU: `model.to(device)`, `data.to(device)`
4236. **Anti-pattern — ignoring batch vector**: In graph classification, always use `global_mean_pool(x, batch)` — forgetting `batch` pools across all graphs
424
425## Common Recipes
426
427### Recipe: Model Explainability (GNNExplainer)
428
429```python
430from torch_geometric.explain import Explainer, GNNExplainer
431
432explainer = Explainer(
433 model=model,
434 algorithm=GNNExplainer(epochs=200),
435 explanation_type='model',
436 node_mask_type='attributes',
437 edge_mask_type='object',
438 model_config=dict(mode='multiclass_classification', task_level='node', return_type='log_probs'),
439)
440
441explanation = explainer(data.x, data.edge_index, index=10)
442print(f'Important edges: {explanation.edge_mask.topk(5).indices}')
443print(f'Important features: {explanation.node_mask[10].topk(5).indices}')
444```
445
446### Recipe: Custom InMemoryDataset
447
448```python
449from torch_geometric.data import InMemoryDataset, Data
450
451class MyDataset(InMemoryDataset):
452 def __init__(self, root, transform=None, pre_transform=None):
453 super().__init__(root, transform, pre_transform)
454 self.load(self.processed_paths[0])
455
456 @property
457 def raw_file_names(self):
458 return ['data.csv']
459
460 @property
461 def processed_file_names(self):
462 return ['data.pt']
463
464 def process(self):
465 data_list = []
466 # Build Data objects from raw files
467 edge_index = torch.tensor([[0, 1], [1, 0]], dtype=torch.long)
468 x = torch.randn(2, 16)
469 data_list.append(Data(x=x, edge_index=edge_index, y=torch.tensor([0])))
470
471 if self.pre_filter is not None:
472 data_list = [d for d in data_list if self.pre_filter(d)]
473 if self.pre_transform is not None:
474 data_list = [self.pre_transform(d) for d in data_list]
475 self.save(data_list, self.processed_paths[0])
476```
477
478### Recipe: Deep GNN with JumpingKnowledge
479
480```python
481from torch_geometric.nn import GCNConv, JumpingKnowledge, LayerNorm
482
483class DeepGNN(torch.nn.Module):
484 def __init__(self, in_ch, hidden, num_layers, out_ch):
485 super().__init__()
486 self.convs = torch.nn.ModuleList()
487 self.norms = torch.nn.ModuleList()
488 self.convs.append(GCNConv(in_ch, hidden))
489 self.norms.append(LayerNorm(hidden))
490 for _ in range(num_layers - 2):
491 self.convs.append(GCNConv(hidden, hidden))
492 self.norms.append(LayerNorm(hidden))
493 self.convs.append(GCNConv(hidden, hidden))
494 self.jk = JumpingKnowledge(mode='cat')
495 self.lin = torch.nn.Linear(hidden * num_layers, out_ch)
496
497 def forward(self, x, edge_index, batch):
498 xs = []
499 for conv, norm in zip(self.convs[:-1], self.norms):
500 x = F.relu(norm(conv(x, edge_index)))
501 xs.append(x)
502 xs.append(self.convs[-1](x, edge_index))
503 return self.lin(global_mean_pool(self.jk(xs), batch))
504```
505
506## Troubleshooting
507
508| Problem | Cause | Solution |
509|---------|-------|---------|
510| `edge_index` shape error | Wrong format (should be [2, E]) | Ensure COO format: `torch.tensor([[src...],[dst...]], dtype=torch.long)` |
511| OOM on large graph | Full-graph forward pass | Use `NeighborLoader` for mini-batch training |
512| Low accuracy | Oversmoothing (too many layers) | Reduce to 2-3 layers, add `JumpingKnowledge` or `PairNorm` |
513| NaN in training | Exploding gradients | Add gradient clipping, reduce learning rate, check feature scale |
514| Wrong graph-level output | Missing `batch` in pooling | Pass `batch` tensor to `global_mean_pool(x, batch)` |
515| Heterogeneous type error | Mismatched node/edge types | Check `data.metadata()` matches model definition |
516| Slow DataLoader | Large graph, no sampling | Use `NeighborLoader` with reasonable `num_neighbors` (e.g., [25,10]) |
517| `x` dimension mismatch | Multi-head attention output | For GATConv: output is `heads*out_channels` unless `concat=False` |
518| Import error for sparse ops | Missing optional dependencies | Install `torch_scatter`, `torch_sparse` from PyG wheels |
519| Pre-transform not applied | Dataset already processed | Delete `processed/` directory and reload |
520
521## Bundled Resources
522
523- **`references/layers_transforms_reference.md`** — Complete catalog of 40+ convolutional layers (GCN, GAT, SAGE, GIN, molecular layers, hypergraph), aggregation operators, pooling (global + hierarchical), normalization layers, pre-built models, auto-encoders, knowledge graph embeddings, utility layers. Transform catalog: structure, feature, spatial, augmentation, mesh, specialized. Consolidated from original layers_reference.md (486 lines) + transforms_reference.md (680 lines). Script functionality (benchmark_model.py, create_gnn_template.py, visualize_graph.py) covered by Core API code blocks and Common Recipes
524- **`references/datasets_catalog.md`** — Comprehensive dataset catalog organized by domain: citation networks (Planetoid, Coauthor, Amazon), graph classification (TUDataset 120+ benchmarks), molecular (QM9, ZINC, MoleculeNet), social (Reddit, Twitch), knowledge graphs (WordNet, FB15k), heterogeneous (OGB_MAG, MovieLens, DBLP), temporal (JODIE), 3D meshes (ShapeNet, ModelNet), OGB integration. Consolidated from original datasets_reference.md (575 lines)
525
526## Related Skills
527
528- **matplotlib-scientific-plotting** — Visualize graph structures, training curves, attention weights
529
530## References
531
532- PyTorch Geometric Documentation: https://pytorch-geometric.readthedocs.io/
533- PyG GitHub: https://github.com/pyg-team/pytorch_geometric
534- Fey & Lenssen (2019), "Fast Graph Representation Learning with PyTorch Geometric", ICLR Workshop