Explaining GNN Predictions: Saliency and Integrated Gradients
Graph neural networks are capable learners, but their decision process is opaque. When a GNN predicts that a molecule is toxic or that a transaction is fraudulent, understanding why that prediction was made is essential for trust, debugging, and scientific discovery. The field of GNN explainability has developed several methods that attribute model predictions back to input features and graph edges. TGraphX implements three of these in the tgraphx.explain module: gradient-based saliency maps, integrated gradients, and edge attribution.
This article explains the theory behind each method, shows practical usage with TGraphX, and discusses the significant limitations of current explainability approaches for GNNs.
Why GNN Explainability Matters
A GNN that classifies drug molecules needs explainability for several reasons. A chemist reviewing a prediction wants to know which functional groups the model found informative. A regulatory submission may require justification for model decisions. A debugging session after unexpected predictions requires understanding whether the model learned meaningful chemistry or spurious correlations.
Beyond deployment, explainability is a research tool. Comparing what a model attends to against known scientific knowledge helps validate whether the model has learned meaningful representations or is exploiting dataset artifacts.
Prerequisites
This article assumes:
- A trained PyTorch model with graph inputs
- Familiarity with TGraphX's Graph and basic GNN layer usage
- Understanding that explainability methods compute input attributions, not natural language explanations
For background on TGraphX model training, see the graph classification architectures tutorial.
Method 1: Gradient-Based Saliency Maps
Saliency maps are the simplest attribution method. The gradient of the output (with respect to a specific class) is computed with respect to the input node features. A high gradient magnitude at node i means that small changes to node i's features would significantly change the prediction, indicating that this node is important to the model.
The mathematical form is:
Saliency(i) = || ∂ f_c(x) / ∂ x_i ||
Where f_c is the model output for class c and x_i is the feature vector of node i. The norm makes the saliency a scalar per node regardless of feature dimensionality.
import torch
from tgraphx import Graph
from tgraphx.explain.saliency import compute_saliency
# Trained model (replace with your actual trained model)
model.eval()
g = Graph(
node_features=torch.randn(50, 16),
edge_index=torch.randint(0, 50, (2, 150), dtype=torch.long),
)
# Compute saliency for class 1 prediction
saliency_scores = compute_saliency(
model=model,
graph=g,
target_class=1,
batch=None, # None for single graph
)
print(saliency_scores.shape) # [50] — one score per node
print(f"Most salient nodes: {saliency_scores.topk(5).indices.tolist()}")
Nodes with high saliency scores are those whose features had the largest gradient magnitude with respect to the predicted class.
Visualizing Saliency
Saliency scores can be normalized and used to color nodes in a graph visualization:
import torch
import matplotlib.pyplot as plt
import networkx as nx
# Normalize to [0, 1]
normed = (saliency_scores - saliency_scores.min()) / (
saliency_scores.max() - saliency_scores.min() + 1e-8
)
# Build a NetworkX graph for visualization
edge_index = g.edge_index.numpy()
G_nx = nx.Graph()
G_nx.add_nodes_from(range(g.node_features.shape[0]))
G_nx.add_edges_from(zip(edge_index[0], edge_index[1]))
pos = nx.spring_layout(G_nx, seed=42)
node_colors = normed.detach().numpy()
plt.figure(figsize=(8, 6))
nx.draw(G_nx, pos=pos, node_color=node_colors, cmap=plt.cm.Reds,
node_size=200, with_labels=True)
plt.title("Node Saliency Map")
plt.colorbar(plt.cm.ScalarMappable(cmap=plt.cm.Reds), label="Saliency")
plt.savefig("saliency_map.png", dpi=150, bbox_inches="tight")
Method 2: Integrated Gradients
Integrated gradients (Sundararajan et al., 2017) address a key failure mode of plain saliency: gradient saturation. In saturated regions of the model (where a large input produces nearly the same output as a slightly larger input), gradients are near zero even though the feature is clearly important. Plain saliency would assign near-zero attribution to a saturated feature, which is misleading.
Integrated gradients compute attributions by integrating gradients along a straight path from a baseline input x' (typically zeros) to the actual input x:
IG_i(x) = (x_i - x'_i) * ∫₀¹ ∂ f(x' + α(x - x')) / ∂ x_i dα
In practice, the integral is approximated with a sum over a fixed number of interpolation steps:
from tgraphx.explain.integrated_gradients import compute_integrated_gradients
ig_scores = compute_integrated_gradients(
model=model,
graph=g,
target_class=1,
baseline="zeros", # baseline input: all-zero node features
num_steps=50, # number of interpolation steps (50 is typical)
batch=None,
)
print(ig_scores.shape) # [50, 16] — attribution per node per feature dimension
# Or reduce to per-node importance:
node_importance = ig_scores.abs().sum(dim=1) # [50]
Integrated gradients satisfy the completeness axiom: the sum of all attributions equals the difference in model output between the baseline and the input. This is a meaningful sanity check that plain saliency does not provide.
Method 3: Edge Attribution
While node attribution asks "which nodes matter?", edge attribution asks "which edges matter?" This is particularly relevant for molecular graphs where specific bonds explain reactivity, or for social graphs where specific relationships explain community membership.
TGraphX's tgraphx.explain.edge_attribution module computes edge importance by masking edges and measuring the change in output:
from tgraphx.explain.edge_attribution import compute_edge_attribution
edge_importance = compute_edge_attribution(
model=model,
graph=g,
target_class=1,
method="gradient", # or "occlusion"
batch=None,
)
print(edge_importance.shape) # [num_edges] — one score per edge
# Find most important edges
top_edge_indices = edge_importance.topk(10).indices
top_edges = g.edge_index[:, top_edge_indices]
print("Top 10 most important edges (source, target):")
for src, tgt in top_edges.T.tolist():
print(f" ({src}, {tgt}): importance = {edge_importance[top_edge_indices].max():.4f}")
The gradient method computes the gradient of the output with respect to each edge's influence. The occlusion method removes edges one at a time and measures the output drop — this is more reliable but O(num_edges) forward passes, making it slow for dense graphs.
Comparing Methods on the Same Prediction
A useful debugging practice is to run multiple attribution methods and check for agreement. When methods agree on the most important nodes or edges, the attribution is more trustworthy. When they disagree strongly, there is ambiguity that deserves investigation:
from tgraphx.explain.saliency import compute_saliency
from tgraphx.explain.integrated_gradients import compute_integrated_gradients
sal = compute_saliency(model, g, target_class=1)
ig = compute_integrated_gradients(model, g, target_class=1, num_steps=50)
sal_ranking = sal.argsort(descending=True)[:10]
ig_ranking = ig.abs().sum(1).argsort(descending=True)[:10]
overlap = set(sal_ranking.tolist()) & set(ig_ranking.tolist())
print(f"Top-10 overlap between Saliency and IG: {len(overlap)}/10")
If overlap is low (fewer than 4-5 nodes), the two methods are identifying different features as important. This is a signal to inspect the model more carefully before trusting either attribution.
Limitations and Honest Notes
GNN explainability is an active research area with significant unresolved problems. The methods in TGraphX are implementations of established techniques, but all carry important caveats.
Gradient-based saliency is not faithful for nonlinear models. Gradients measure local sensitivity, not global importance. A node that is consistently important throughout training may have near-zero gradient at evaluation time because the model has learned to handle it deterministically. Saliency maps reflect the local model landscape at inference, not the global training signal.
Integrated gradients depend heavily on the baseline choice. With a zero baseline, IG attributes importance relative to "no feature information." With a random baseline or a mean baseline, the attributions can change substantially. There is no universally correct baseline, and different choices lead to different attributions that can all satisfy the completeness axiom.
Edge attribution via occlusion is computationally expensive. For a graph with 10,000 edges, computing occlusion-based edge importance requires 10,001 forward passes. Use the gradient-based edge method for large graphs.
None of these methods produce human-interpretable explanations for non-expert users. A list of important node indices or edge weights requires domain knowledge to interpret. For a toxicity prediction, you need a chemist to translate "nodes 4, 7, 12 are most important" into a chemical substructure explanation.
Over-smoothed GNNs produce meaningless attributions. When a GNN has been over-trained or uses too many layers, all node representations converge to similar values. Attribution methods will then assign uniform or near-uniform importance scores across all nodes, which is technically correct (all nodes are equally unimportant when they all look the same) but provides no useful information.
TGraphX explainability is experimental. The explain module is newer than the core GNN layers and has been validated on fewer configurations. Verify that attribution scores behave sensibly on toy examples before applying them to production models.
A Minimal End-to-End Example
import torch
import torch.nn.functional as F
from tgraphx import Graph
from tgraphx.layers.gin import TensorGINLayer
from tgraphx.layers.pooling import GlobalSumPool
from tgraphx.explain.saliency import compute_saliency
import torch.nn as nn
class SmallGNN(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = TensorGINLayer(8, 16, train_eps=True)
self.conv2 = TensorGINLayer(16, 32, train_eps=True)
self.pool = GlobalSumPool()
self.fc = nn.Linear(32, 2)
def forward(self, x, edge_index, batch=None):
x = F.relu(self.conv1(x, edge_index))
x = F.relu(self.conv2(x, edge_index))
if batch is not None:
x = self.pool(x, batch)
else:
x = x.mean(0, keepdim=True)
return self.fc(x)
model = SmallGNN()
# In practice, load trained weights: model.load_state_dict(torch.load("model.pt"))
g = Graph(
node_features=torch.randn(20, 8),
edge_index=torch.randint(0, 20, (2, 60), dtype=torch.long),
)
saliency = compute_saliency(model, g, target_class=0)
top_nodes = saliency.topk(3).indices.tolist()
print(f"Top 3 nodes by saliency: {top_nodes}")
For related methods that approach interpretability from the graph structure side, see the graph mining overview. For ensuring that models being explained were trained reproducibly, see the GNN research reproducibility guide.
Applying Explainability to Node Classification
The examples above focus on graph classification. For node classification tasks, the workflow is similar but the target is a specific node's prediction rather than a graph-level prediction:
from tgraphx.explain.saliency import compute_node_saliency
# Compute saliency for a specific node's prediction
target_node = 42
target_class = 1
node_saliency = compute_node_saliency(
model=model,
graph=g,
node_idx=target_node,
target_class=target_class,
)
print(f"Node {target_node} prediction attribution:")
print(f" Self saliency: {node_saliency[target_node]:.4f}")
print(f" Neighbor contributions: {node_saliency[g.edge_index[1][g.edge_index[0] == target_node]]}")
In node classification, the neighbors of the target node are typically the most informative in the attribution — they provided the context for the prediction through message passing. Edges to high-saliency neighbors are likely the most important structural signal.
Explainability and Trust in Research Papers
Using explainability methods in research papers requires caution. Attribution maps are often presented as ground truth in GNN papers, but they are heuristics subject to the limitations described above. Specific guidance for research use:
Report the method and its parameters. "Gradient saliency with ReLU non-linearity" and "Integrated Gradients with 50 steps and zero baseline" produce different results even on the same model. Enough detail must be given for readers to reproduce your attribution plots.
Validate attributions on toy examples. Before presenting explanations on real data, verify that the method correctly attributes known-important features on a controlled synthetic example where the right answer is known.
Do not conflate explanation quality with model quality. A model can be accurate but poorly explained, or produce plausible-looking but unfaithful explanations. Attribution quality does not imply prediction quality.
Consider attribution consistency across seeds. If two independently trained models with the same architecture agree on the most important nodes/edges, the explanation is more likely to reflect genuine learned structure than model-specific artifacts. Running the same attribution on three or more models and looking for overlap is a cheap robustness check.
What This Article Builds On
The attribution methods covered here operate on trained PyTorch models. Training a TGraphX model is covered in the MNIST as a graph tutorial and the shape-aware validation guide. The over-smoothing issue mentioned in the limitations section — which makes attribution meaningless for deep GNNs — is discussed in depth in the articles hub.
For the graph mining perspective on structural importance (centrality, motifs), which provides complementary non-gradient-based explanations, see the graph mining overview.
Frequently Asked Questions
Can I use these methods with PyTorch Geometric models? Yes, with caveats. The TGraphX explain module expects tgraphx.Graph objects. If your model takes PyG Data objects, you will need to adapt the wrapper. The underlying gradient computation is standard PyTorch autograd and is model-agnostic.
Do saliency methods work for tensor-valued node features? Yes. compute_saliency returns gradients with respect to the full input tensor, so for [N, C, H, W] node features, the saliency for each node is a [C, H, W] tensor. You can reduce it to a scalar by taking the Frobenius norm or the max absolute value.
How do I choose between saliency and integrated gradients? Use saliency for fast iteration and debugging. Use integrated gradients when you need to report results in a paper or when saturation is suspected (e.g., heavily regularized models or models trained with aggressive clipping).
Is GNNExplainer or SubgraphX available in TGraphX? Not in the current release. The explain module covers gradient-based methods. GNNExplainer (which learns a mask via optimization) and SubgraphX (which uses Monte Carlo tree search) are more complex and are not currently in scope.