TGraphX › Insights › Graph-Level Classification: Architectures and Pooling
← Back to Insights

Graph-Level Classification: Architectures and Pooling

Target keyword: graph classification GNN pytorch

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:

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

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

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

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

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

python
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

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

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

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