Message Passing Neural Networks
Message Passing Neural Networks let graph nodes iteratively exchange vector messages with immediate neighbors to update their internal representations.
Why Does This Exist?
Between 2015 and 2017, the machine learning literature was flooded with competing graph deep learning architectures: Graph Convolutional Networks (Kipf & Welling), GraphSAGE (Hamilton et al.), Molecular Graph Convolutions (Duvenaud et al.), and Gated Graph Neural Networks (Li et al.). Each used different terminology, notation, and tensor operations, making it difficult to understand whether these models were fundamentally distinct architectures or minor variations of the same underlying principle.
Gilmer et al. (2017) cut through this fragmentation by introducing the Message Passing Neural Network (MPNN) framework. MPNN demonstrated that nearly all spatial graph neural networks are instances of a single elegant computational paradigm: an iterative message exchange between neighboring nodes. Furthermore, MPNN proved that by formalizing these operations, one can guarantee permutation equivariance (reordering node indices produces an identically reordered output) and permutation invariance (predicting a whole-graph property like molecular solubility is completely independent of the arbitrary order in which nodes are indexed).
Think of It Like This
A round-table discussion among department liaisons
Imagine an organization where team managers sit in individual offices connected by telephone lines. Each manager holds a private notepad summarizing their department's current status.
At the top of the hour, every manager executes three steps:
- Compose Messages: For every connected colleague, the manager drafts a targeted briefing memo based on what they know and the specific nature of their partnership.
- Collect Inboxes: Each manager receives the memos delivered to their door. Because memos arrive simultaneously, the manager stacks them into an inbox without caring which courier dropped off their letter first. They tally or average the figures across all memos.
- Update Notes: The manager reads the aggregated briefing summary, compares it against their own previous status notepad, and writes a revised status entry.
After one round of calls, every manager knows what their immediate neighbors know (1-hop awareness). After two rounds, news from two offices away has diffused across the organization (2-hop awareness).
How It Actually Works
The Three Phases: Message, Aggregate, and Update
An MPNN operates on a graph where nodes possess hidden states and edges possess optional feature attributes . The forward pass repeats over discrete message-passing time steps.
Each step consists of three algebraic phases:
1. Message Phase
For every directed edge , a parameterized message function computes an outgoing message vector from the source node state , destination state , and edge attributes :
In Graph Convolutional Networks (GCN), is a scaled linear projection: . In Graph Attention Networks (GAT), .
2. Aggregation Phase
Node gathers incoming messages from all adjacent nodes in its neighborhood using a permutation-invariant reduction operator :
Because satisfies commutativity () and associativity (), the aggregated message is strictly invariant to node permutation.
3. Update Phase
The updated node representation is computed by an update function combining the node's prior state with the aggregated incoming message:
Common choices for include a Gated Recurrent Unit (GRU cell in GGNN) or a Multi-Layer Perceptron with a residual connection: .
4. Readout Phase
For graph-level prediction tasks (e.g., predicting chemical toxicity), a permutation-invariant readout function pools all final representations :
Worked Example
Let node have two neighbors and . Features are 2-dimensional ().
- Initial states:
- Linear message weight matrix :
- Aggregator: (Sum).
- Update function: with .
-
Compute Messages:
-
Aggregate Messages:
-
Update State: Node incorporated structural information from both neighbors while retaining its identity.
Code
from typing import List, Tupleimport numpy as np
class MPNNLayer: """Minimal single-step Message Passing Neural Network layer."""
def __init__(self, in_dim: int, out_dim: int) -> None: rng = np.random.default_rng(seed=42) # Message transform matrix self.W_msg: np.ndarray = rng.standard_normal((in_dim, out_dim)) * 0.2 # Self-update transform matrix self.W_self: np.ndarray = rng.standard_normal((in_dim, out_dim)) * 0.2
def forward( self, node_features: np.ndarray, # Shape: (N, in_dim) edge_index: List[Tuple[int, int]], # List of directed (src, dst) edges ) -> np.ndarray: num_nodes = node_features.shape[0] out_dim = self.W_msg.shape[1]
# 1. Message Phase: Compute message from each src to dst # 2. Aggregation Phase: Sum into accumulator aggregated_messages = np.zeros((num_nodes, out_dim), dtype=np.float64) for src, dst in edge_index: msg = np.dot(node_features[src], self.W_msg) aggregated_messages[dst] += msg
# 3. Update Phase: Combine self state with incoming messages + ReLU self_transformed = np.dot(node_features, self.W_self) updated_states = np.maximum(0.0, self_transformed + aggregated_messages) return updated_states
# Graph: Triangle where 0 <-> 1, 1 <-> 2, 2 <-> 0edges: List[Tuple[int, int]] = [ (0, 1), (1, 0), (1, 2), (2, 1), (2, 0), (0, 2),]features = np.array([ [1.0, 0.0], [0.5, 1.0], [-0.5, 0.5],], dtype=np.float64)
layer = MPNNLayer(in_dim=2, out_dim=2)out_h = layer.forward(features, edges)print("Updated node embeddings shape:", out_h.shape)# -> Updated node embeddings shape: (3, 2)print("Node 0 new state:", np.round(out_h[0], 4))Watch Out For
The 1-Weisfeiler-Lehman expressivity bottleneck
Standard MPNN message passing with anonymous node aggregations is theoretically bounded in expressive power by the 1-Weisfeiler-Lehman (1-WL) graph isomorphism test (Morris et al., 2019). An MPNN cannot distinguish between strongly regular graphs that share identical degree distributions and local neighbor multiset structures, such as two disconnected triangles versus one 6-cycle ring.
If your problem domain requires detecting specific molecular substructures (like benzene rings, cliques, or cycles) that 1-WL fails to separate, do not simply stack more MPNN layers. Incorporate structural cycle features, relative random walk positional encodings, or higher-order Subgraph/Cellular GNN architectures.
The Quick Version
- MPNN provides a unifying three-stage framework for spatial graph neural networks: Message computation, Aggregation, and State Update.
- Permutation equivariance is guaranteed by selecting commutative and associative aggregation operators such as sum, mean, or max.
- The theoretical expressive power of standard MPNNs is bounded by the 1-Weisfeiler-Lehman graph isomorphism test.