Graph-Level Classification: Architectures and Pooling
Node classification asks about individual nodes. Graph classification asks about entire graphs: given a collection of graphs, predict a label for each one. The two tasks require fundamentally different model architectures. For node classification, the GNN reads out a per-node representation. For graph classification, those per-node representations must be aggregated into a single vector before the final classifier. The choice of aggregation — pooling — has a large effect on what structural information is preserved.
This article covers the graph classification task, the major architecture patterns (GCN+pool, GIN+pool), global versus hierarchical pooling, how TGraphX's GraphBatch handles multiple graphs, and practical code for a graph classifier.
The Graph Classification Setup
In graph classification, the dataset is a collection of graphs {G_1, G_2, ..., G_N}, each with a label y_i. The labels might be:
- Molecule property (toxic/non-toxic, soluble/insoluble)
- Protein function category
- Social network type (e.g., communication vs. collaboration)
- Program behavior class (benign vs. malicious)
The challenge is that graphs in the dataset have different sizes: G_1 might have 20 nodes and G_2 might have 400. A fixed-size classifier cannot directly accept variable-size inputs, so pooling is necessary.
Prerequisites
This article assumes familiarity with:
- Basic GNN operations (message passing, aggregation)
- PyTorch module structure
- The TGraphX Graph and GraphBatch objects
See the MNIST as a graph tutorial for an introduction to TGraphX's core data structures.
Global Pooling
Global pooling maps a set of node representations {h_1, ..., h_n} to a single vector. TGraphX provides three variants via tgraphx.layers.pooling:
import torch
from tgraphx.layers.pooling import GlobalMeanPool, GlobalSumPool, GlobalMaxPool
# Suppose we have node representations from a GNN
# x: [N, D] — node features after message passing
# batch: [N] — which graph each node belongs to (for batched processing)
x = torch.randn(100, 64) # 100 nodes, 64-dim features
batch = torch.cat([torch.full((20,), i) for i in range(5)]) # 5 graphs of 20 nodes
mean_pool = GlobalMeanPool()
sum_pool = GlobalSumPool()
max_pool = GlobalMaxPool()
graph_repr_mean = mean_pool(x, batch) # [5, 64]
graph_repr_sum = sum_pool(x, batch) # [5, 64]
graph_repr_max = max_pool(x, batch) # [5, 64]
GlobalMeanPool averages all node representations. It is invariant to graph size and works well when all nodes contribute roughly equally to the graph label. It is the most common default choice.
GlobalSumPool sums node representations. It is size-sensitive: larger graphs tend to produce larger sum vectors. This is actually a theoretical advantage for tasks where graph size correlates with the label (e.g., molecule size and toxicity).
GlobalMaxPool takes the element-wise maximum. It selects the most activated feature value across all nodes, effectively asking "does any node in this graph have high activation in this dimension?" This is useful when the label depends on the presence of a specific local structure rather than the overall composition.
Architecture 1: GCN + Global Pool
The simplest graph classification architecture applies two GCN-style message passing layers, then global mean pooling:
import torch
import torch.nn as nn
import torch.nn.functional as F
from tgraphx.layers.sage import TensorGraphSAGELayer
from tgraphx.layers.pooling import GlobalMeanPool
class SimpleGraphClassifier(nn.Module):
def __init__(self, in_dim, hidden_dim, out_dim, num_classes):
super().__init__()
self.conv1 = TensorGraphSAGELayer(in_dim, hidden_dim)
self.conv2 = TensorGraphSAGELayer(hidden_dim, out_dim)
self.pool = GlobalMeanPool()
self.classifier = nn.Linear(out_dim, num_classes)
def forward(self, x, edge_index, batch):
x = F.relu(self.conv1(x, edge_index))
x = F.relu(self.conv2(x, edge_index))
x = self.pool(x, batch) # [B, out_dim]
return self.classifier(x) # [B, num_classes]
This architecture is fast and works well on many benchmarks. Its weakness is that the global pooling step discards structural information about which nodes were activated and how they relate to each other.
Architecture 2: GIN + Sum Pool
For graph classification tasks where distinguishing non-isomorphic graphs matters, GIN with sum pooling is theoretically stronger:
from tgraphx.layers.gin import TensorGINLayer
from tgraphx.layers.pooling import GlobalSumPool
class GINGraphClassifier(nn.Module):
def __init__(self, in_dim, hidden_dim, out_dim, num_classes, num_layers=3):
super().__init__()
self.convs = nn.ModuleList()
self.bns = nn.ModuleList()
dims = [in_dim] + [hidden_dim] * (num_layers - 1) + [out_dim]
for i in range(num_layers):
self.convs.append(
TensorGINLayer(dims[i], dims[i+1], train_eps=True, use_batchnorm=False)
)
self.bns.append(nn.BatchNorm1d(dims[i+1]))
self.pool = GlobalSumPool()
self.classifier = nn.Sequential(
nn.Linear(out_dim, hidden_dim),
nn.ReLU(),
nn.Dropout(0.5),
nn.Linear(hidden_dim, num_classes),
)
def forward(self, x, edge_index, batch):
for conv, bn in zip(self.convs, self.bns):
x = conv(x, edge_index)
x = bn(x)
x = F.relu(x)
x = self.pool(x, batch)
return self.classifier(x)
The multi-layer readout (concatenating representations from each layer before pooling) is a common extension of GIN for graph classification:
def forward(self, x, edge_index, batch):
xs = []
for conv, bn in zip(self.convs, self.bns):
x = F.relu(bn(conv(x, edge_index)))
xs.append(self.pool(x, batch))
# Concatenate pooled representations from all layers
return self.classifier(torch.cat(xs, dim=1))
Using GraphBatch for Multiple Graphs
The GraphBatch object in TGraphX handles batching of variable-size graphs efficiently. It concatenates node features and adjusts edge indices so all graphs appear as a single large disconnected graph:
from tgraphx import Graph, GraphBatch
# Create individual graphs
graphs = [
Graph(
node_features=torch.randn(20, 16),
edge_index=torch.randint(0, 20, (2, 60), dtype=torch.long),
graph_label=torch.tensor([0]),
),
Graph(
node_features=torch.randn(35, 16),
edge_index=torch.randint(0, 35, (2, 90), dtype=torch.long),
graph_label=torch.tensor([1]),
),
Graph(
node_features=torch.randn(15, 16),
edge_index=torch.randint(0, 15, (2, 40), dtype=torch.long),
graph_label=torch.tensor([0]),
),
]
# Batch them together
batch = GraphBatch.from_graphs(graphs)
print(f"Total nodes: {batch.node_features.shape[0]}") # 70
print(f"Batch vector: {batch.batch.shape}") # [70]
print(f"Labels: {batch.graph_labels}") # [0, 1, 0]
The batch.batch vector contains the graph index for each node, which is exactly what the pooling layers need:
model = GINGraphClassifier(in_dim=16, hidden_dim=32, out_dim=64, num_classes=2)
logits = model(batch.node_features, batch.edge_index, batch.batch)
print(logits.shape) # [3, 2]
Training a Graph Classifier
from torch.utils.data import DataLoader
import torch.optim as optim
# Assume graph_dataset is a list of (Graph, label) tuples
def collate_graphs(batch_list):
graphs = [item[0] for item in batch_list]
return GraphBatch.from_graphs(graphs)
loader = DataLoader(
graph_dataset, batch_size=32, shuffle=True, collate_fn=collate_graphs
)
model = GINGraphClassifier(in_dim=16, hidden_dim=64, out_dim=128, num_classes=5)
optimizer = optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4)
criterion = nn.CrossEntropyLoss()
for epoch in range(100):
model.train()
total_loss = 0.0
for batched_graphs in loader:
optimizer.zero_grad()
logits = model(
batched_graphs.node_features,
batched_graphs.edge_index,
batched_graphs.batch,
)
loss = criterion(logits, batched_graphs.graph_labels)
loss.backward()
optimizer.step()
total_loss += loss.item()
if epoch % 10 == 0:
print(f"Epoch {epoch}: loss = {total_loss / len(loader):.4f}")
Architecture Comparison
| Architecture | Pooling | Expressiveness | Key Strength | Key Weakness |
|---|---|---|---|---|
| SAGE + GlobalMean | Global mean | Low-medium | Fast, simple, size-invariant | Loses structural detail |
| SAGE + GlobalSum | Global sum | Low-medium | Size-sensitive outputs | Sum can explode for large graphs |
| GIN + GlobalSum | Global sum | WL-equivalent | Maximum expressiveness | More parameters, slower |
| GIN + Hierarchical (future) | SAGPool/DiffPool | High | Preserves local structure | Complex, harder to tune |
| Multi-layer readout GIN | Sum per layer + concat | WL-equivalent | Multi-scale representation | Higher memory use |
The "best" architecture depends on the dataset. On molecular property prediction benchmarks, GIN with sum pooling is a strong baseline. On social network graphs with high-degree nodes, mean pooling often performs comparably at lower compute cost.
Limitations and Honest Notes
Global pooling discards positional information. Two graphs with the same multiset of node features but different edge structure may produce identical representations after global mean pooling. GIN's injective aggregation mitigates this at the message-passing stage, but global pooling is still a lossy final step.
Hierarchical pooling methods (DiffPool, SAGPool, MinCutPool) are not in the current TGraphX release. These methods learn to progressively coarsen the graph, preserving structural information through multiple abstraction levels. They are significantly more complex to implement and train. The global pooling methods documented here are the current TGraphX offering.
Batch normalization with small batches causes instability. If your batch contains very few graphs, per-node batch normalization estimates are noisy. Consider using layer normalization or larger batch sizes.
GraphBatch edge index adjustment assumes node indices are local per graph. If you construct edge indices with global node IDs rather than per-graph IDs, GraphBatch.from_graphs will not produce correct batching. Always use 0-indexed per-graph node IDs.
For graph-level feature analysis tools, see the graph mining overview. For shape-aware validation to catch feature tensor mismatches early, see the shape-aware validation guide.
Evaluation Protocol for Graph Classification
Graph classification evaluation requires careful attention to splitting. The standard protocol holds out 10-20% of graphs as a test set:
from sklearn.model_selection import StratifiedKFold
import numpy as np
# 10-fold cross-validation (standard for molecular benchmarks)
labels = [g.graph_label.item() for g in graph_dataset]
kfold = StratifiedKFold(n_splits=10, shuffle=True, random_state=42)
fold_accs = []
for fold, (train_idx, test_idx) in enumerate(kfold.split(graph_dataset, labels)):
train_graphs = [graph_dataset[i] for i in train_idx]
test_graphs = [graph_dataset[i] for i in test_idx]
# Train model on train_graphs, evaluate on test_graphs
acc = train_and_evaluate(train_graphs, test_graphs)
fold_accs.append(acc)
print(f"Fold {fold+1}: accuracy = {acc:.4f}")
print(f"Mean accuracy: {np.mean(fold_accs):.4f} ± {np.std(fold_accs):.4f}")
10-fold cross-validation is standard for datasets like TU benchmark datasets (MUTAG, PROTEINS, IMDB-B). Reporting mean and standard deviation across folds is required — reporting only the best fold is misleading. For reproducibility, always report the random seed and cross-validation strategy.
Handling Graphs Without Node Features
Some graph datasets have no node features — only topology. A common strategy is to assign each node a degree-based or constant feature:
import torch
from tgraphx import Graph
def add_degree_features(g: Graph) -> Graph:
"""Replace empty or missing node features with node degree as a feature."""
N = g.edge_index.max().item() + 1
degree = torch.zeros(N)
for src, dst in g.edge_index.T.tolist():
degree[src] += 1
degree[dst] += 1
degree_feature = degree.unsqueeze(1) # [N, 1]
return Graph(
node_features=degree_feature,
edge_index=g.edge_index,
graph_label=g.graph_label,
)
For more expressive features without labeled data, Weisfeiler-Lehman (WL) color refinement features or positional encodings can be computed from the graph topology using tgraphx.mining.structural.
Connecting Graph Classification to Research
Graph classification is one of the primary benchmarks for evaluating GNN expressiveness. The question of whether two GNNs can be distinguished by their graph classification performance is directly related to the Weisfeiler-Lehman hierarchy. GIN, being WL-equivalent, is theoretically the most expressive standard message-passing GNN for this task.
In practice, differences between GCN, SAGE, GAT, and GIN on graph classification benchmarks are often smaller than the variance across random seeds. This is a well-documented phenomenon in the GNN literature: the choice of architecture matters less than the choice of dataset preprocessing, normalization strategy, and training recipe. See the GNN research reproducibility guide for how to report graph classification results honestly.
What This Article Builds On
This article assumes familiarity with TGraphX's core Graph and GraphBatch objects. If you are new to graph construction in TGraphX, start with the MNIST as a graph tutorial, which covers the same Graph API in the node classification setting. The shape-aware validation guide covers validation patterns that are equally important in the graph classification setting.
For the broader context of GNN architectures and when to use each layer type, see the articles hub.