Skip to content
AI360Xpert
Beta

Message Passing Neural Networks

Message Passing Neural Networks let graph nodes iteratively exchange vector messages with immediate neighbors to update their internal representations.

MPNN unifies graph neural networks into three modular phases: message generation along edges, permutation-invariant aggregation, and node state updating
MPNN unifies graph neural networks into three modular phases: message generation along edges, permutation-invariant aggregation, and node state updating

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:

  1. 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.
  2. 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.
  3. 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 G=(V,E)G = (V, E) where nodes possess hidden states hv∈Rdv\mathbf{h}_v \in \mathbb{R}^{d_v} and edges possess optional feature attributes euv∈Rde\mathbf{e}_{uv} \in \mathbb{R}^{d_e}. The forward pass repeats over TT discrete message-passing time steps.

Each step t∈{1,…,T}t \in \{1, \dots, T\} consists of three algebraic phases:

1. Message Phase

For every directed edge (u,v)∈E(u, v) \in E, a parameterized message function MtM_t computes an outgoing message vector muv(t)\mathbf{m}_{uv}^{(t)} from the source node state hu(t−1)\mathbf{h}_u^{(t-1)}, destination state hv(t−1)\mathbf{h}_v^{(t-1)}, and edge attributes euv\mathbf{e}_{uv}:

muv(t)=Mt(hv(t−1),hu(t−1),euv)\mathbf{m}_{uv}^{(t)} = M_t\left(\mathbf{h}_v^{(t-1)}, \mathbf{h}_u^{(t-1)}, \mathbf{e}_{uv}\right)

In Graph Convolutional Networks (GCN), MtM_t is a scaled linear projection: Mt=1deg⁡(v)deg⁡(u)W(t)huM_t = \frac{1}{\sqrt{\deg(v)\deg(u)}} \mathbf{W}^{(t)} \mathbf{h}_u. In Graph Attention Networks (GAT), Mt=αuv(t)W(t)huM_t = \alpha_{uv}^{(t)} \mathbf{W}^{(t)} \mathbf{h}_u.

2. Aggregation Phase

Node vv gathers incoming messages from all adjacent nodes in its neighborhood N(v)\mathcal{N}(v) using a permutation-invariant reduction operator ⨁∈{∑,mean,max⁡}\bigoplus \in \{\sum, \text{mean}, \max\}:

mv(t)=⨁u∈N(v)muv(t)\mathbf{m}_v^{(t)} = \bigoplus_{u \in \mathcal{N}(v)} \mathbf{m}_{uv}^{(t)}

Because ⨁\bigoplus satisfies commutativity (a⊕b=b⊕aa \oplus b = b \oplus a) and associativity ((a⊕b)⊕c=a⊕(b⊕c)(a \oplus b) \oplus c = a \oplus (b \oplus c)), the aggregated message mv(t)\mathbf{m}_v^{(t)} is strictly invariant to node permutation.

3. Update Phase

The updated node representation hv(t)\mathbf{h}_v^{(t)} is computed by an update function UtU_t combining the node's prior state with the aggregated incoming message:

hv(t)=Ut(hv(t−1),mv(t))\mathbf{h}_v^{(t)} = U_t\left(\mathbf{h}_v^{(t-1)}, \mathbf{m}_v^{(t)}\right)

Common choices for UtU_t include a Gated Recurrent Unit (GRU cell in GGNN) or a Multi-Layer Perceptron with a residual connection: σ(Wselfhv(t−1)+mv(t))\sigma(\mathbf{W}_{\text{self}} \mathbf{h}_v^{(t-1)} + \mathbf{m}_v^{(t)}).

4. Readout Phase

For graph-level prediction tasks (e.g., predicting chemical toxicity), a permutation-invariant readout function RR pools all final representations hv(T)\mathbf{h}_v^{(T)}:

y^=R({hv(T)∣v∈V})=MLP(∑v∈Vσ(Wghv(T))⊙tanh⁡(Wfhv(T)))\hat{\mathbf{y}} = R\left(\{\mathbf{h}_v^{(T)} \mid v \in V\}\right) = \text{MLP}\left( \sum_{v \in V} \sigma(\mathbf{W}_g \mathbf{h}_v^{(T)}) \odot \tanh(\mathbf{W}_f \mathbf{h}_v^{(T)}) \right)

Worked Example

Let node vv have two neighbors u1u_1 and u2u_2. Features are 2-dimensional (d=2d = 2).

  • Initial states: hv(0)=[1.00.0],hu1(0)=[0.51.0],hu2(0)=[−0.50.5]\mathbf{h}_v^{(0)} = \begin{bmatrix} 1.0 \\ 0.0 \end{bmatrix}, \quad \mathbf{h}_{u_1}^{(0)} = \begin{bmatrix} 0.5 \\ 1.0 \end{bmatrix}, \quad \mathbf{h}_{u_2}^{(0)} = \begin{bmatrix} -0.5 \\ 0.5 \end{bmatrix}
  • Linear message weight matrix WM∈R2×2\mathbf{W}_M \in \mathbb{R}^{2 \times 2}: WM=[0.50.00.00.5]\mathbf{W}_M = \begin{bmatrix} 0.5 & 0.0 \\ 0.0 & 0.5 \end{bmatrix}
  • Aggregator: ⨁=∑\bigoplus = \sum (Sum).
  • Update function: U(hv,mv)=ReLU(WUhv+mv)U(\mathbf{h}_v, \mathbf{m}_v) = \text{ReLU}(\mathbf{W}_U \mathbf{h}_v + \mathbf{m}_v) with WU=[1.00.00.01.0]\mathbf{W}_U = \begin{bmatrix} 1.0 & 0.0 \\ 0.0 & 1.0 \end{bmatrix}.
  1. Compute Messages: mu1v(1)=WMhu1(0)=[0.5×0.50.5×1.0]=[0.250.50]\mathbf{m}_{u_1 v}^{(1)} = \mathbf{W}_M \mathbf{h}_{u_1}^{(0)} = \begin{bmatrix} 0.5 \times 0.5 \\ 0.5 \times 1.0 \end{bmatrix} = \begin{bmatrix} 0.25 \\ 0.50 \end{bmatrix} mu2v(1)=WMhu2(0)=[0.5×(−0.5)0.5×0.5]=[−0.250.25]\mathbf{m}_{u_2 v}^{(1)} = \mathbf{W}_M \mathbf{h}_{u_2}^{(0)} = \begin{bmatrix} 0.5 \times (-0.5) \\ 0.5 \times 0.5 \end{bmatrix} = \begin{bmatrix} -0.25 \\ 0.25 \end{bmatrix}

  2. Aggregate Messages: mv(1)=mu1v(1)+mu2v(1)=[0.25+(−0.25)0.50+0.25]=[0.000.75]\mathbf{m}_v^{(1)} = \mathbf{m}_{u_1 v}^{(1)} + \mathbf{m}_{u_2 v}^{(1)} = \begin{bmatrix} 0.25 + (-0.25) \\ 0.50 + 0.25 \end{bmatrix} = \begin{bmatrix} 0.00 \\ 0.75 \end{bmatrix}

  3. Update State: zv=WUhv(0)+mv(1)=[1.00.0]+[0.000.75]=[1.000.75]\mathbf{z}_v = \mathbf{W}_U \mathbf{h}_v^{(0)} + \mathbf{m}_v^{(1)} = \begin{bmatrix} 1.0 \\ 0.0 \end{bmatrix} + \begin{bmatrix} 0.00 \\ 0.75 \end{bmatrix} = \begin{bmatrix} 1.00 \\ 0.75 \end{bmatrix} hv(1)=ReLU([1.000.75])=[1.000.75]\mathbf{h}_v^{(1)} = \text{ReLU}\left(\begin{bmatrix} 1.00 \\ 0.75 \end{bmatrix}\right) = \begin{bmatrix} 1.00 \\ 0.75 \end{bmatrix} Node vv 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.