Graph Neural Networks (GNNs)
Standard Deep Learning requires data to be in strict structures: Grids (images for CNNs) or Sequences (text for Transformers). But what if your data is a messy, interconnected web?
- Social Networks (Users are nodes, friendships are edges)
- Chemistry (Atoms are nodes, bonds are edges)
- Google Maps (Intersections are nodes, roads are edges)
Graph Neural Networks (GNNs) are designed to process this interconnected data.
Message Passing
The core mechanism of a GNN is Message Passing. In a Graph Convolutional Network (GCN), every node updates its own embedding by "listening" to the embeddings of its immediate neighbors.
If you have a 3-layer GNN, each node gathers information from its neighbors, its neighbors' neighbors, and its neighbors' neighbors' neighbors. Over time, each node's vector becomes a rich summary of its local graph neighborhood.
Key Tasks
- Node Classification: E.g., Predicting if a user in a social network is a bot.
- Link Prediction: E.g., Predicting if two users should be friends (used heavily in recommendations!).
- Graph Classification: E.g., Predicting if a molecular graph will be a toxic drug or an effective medicine.
Python Implementation: pytorchGeometric
# PyTorch Geometric (PyG) is the standard library for GNNs
import torch
import torch.nn.functional as F
from torch_geometric.nn import GCNConv
# Define a simple 2-layer Graph Convolutional Network
class GCN(torch.nn.Module):
def __init__(self, num_node_features, num_classes):
super().__init__()
# First Graph Convolution Layer
self.conv1 = GCNConv(num_node_features, 16)
# Second Graph Convolution Layer
self.conv2 = GCNConv(16, num_classes)
def forward(self, data):
x, edge_index = data.x, data.edge_index
# Message Passing Step 1
x = self.conv1(x, edge_index)
x = F.relu(x)
x = F.dropout(x, training=self.training)
# Message Passing Step 2
x = self.conv2(x, edge_index)
return F.log_softmax(x, dim=1)