本文目录导读:

消息传递神经网络(Message Passing Neural Network, MPNN) 是一种用于处理图结构数据(Graph Data)的通用框架,它由 Gilmer 等人于 2017 年在论文 Neural Message Passing for Quantum Chemistry 中提出。
MPNN 的核心思想非常直观:图中的每个节点通过不断与它的邻居节点“交换信息”(消息),来更新自己对于整个图结构的理解。
可以把 MPNN 想象成一个社交网络:
- 每个人(节点)最初只知道自己的一些属性(特征向量)。
- 每个时间步,大家都向自己的朋友(邻居)发送一条描述当前状态的消息。
- 每个人收集所有朋友发来的消息后进行聚合,再结合自己原本的状态进行“思考”(更新),形成新的认知。
- 重复这个过程多次后,每个人的认知就包含了整个朋友圈的结构信息。
MPNN 的核心框架
MPNN 的计算过程通常分为两个阶段:消息传递阶段 和 读出阶段。
消息传递阶段 (Message Passing Phase)
这个阶段会执行 T 步 迭代,对于每一层 t(从 1 到 T):
-
消息函数 (Message Function) $M_t$:对于图中的每条边
(u, v),根据节点u和v在当前时间步t-1的特征,计算从节点u要传递给节点v的“消息”。- $mv^{(t)} = \sum{u \in N(v)} M_t(h_v^{(t-1)}, hu^{(t-1)}, e{vu})$
- $h_v^{(t-1)}$ 是节点
v在上一步的特征。 - $h_u^{(t-1)}$ 是邻居节点
u在上一步的特征。 - $e_{vu}$ 是边
(u,v)的特征(如果有的话,例如化学键类型)。 - $N(v)$ 是节点
v的所有邻居集合。 - 上面公式使用了求和聚合,这是最常见的形式,MPNN 框架允许使用任何置换不变函数(如求和、均值、最大值)来聚合邻居信息。
- $h_v^{(t-1)}$ 是节点
-
更新函数 (Update Function) $U_t$:节点
v将聚合后的消息 $m_v^{(t)}$ 与自身上一层的特征 $h_v^{(t-1)}$ 结合起来,更新自己的特征。$h_v^{(t)} = U_t(h_v^{(t-1)}, m_v^{(t)})$
经过 T 步后,最终的节点特征 $h_v^{(T)}$ 就是节点的“上下文感知”表示,包含了其 T 跳范围内的结构信息。
读出阶段 (Readout Phase)
当所有节点都有了最终的特征表示后,如果需要做图级别的任务(例如预测分子毒性、判断图的类别),就需要一个读出函数 R 将所有节点的特征聚合成一个代表整个图的特征向量 $y$。
- $y = R({h_v^{(T)} | v \in G})$
- 读出函数同样需要是置换不变的(即无论节点顺序如何,结果都一样),例如对所有节点特征求和、取平均、或者使用更复杂的 Set2Set、SortPool 等。
一个具体的例子:GCN (图卷积网络)
GCN 可以看作是 MPNN 的一个经典特例,在 GCN 中:
-
消息函数
- $mv^{(t)} = \sum{u \in N(v) \cup {v}} \frac{1}{\sqrt{deg(v) \cdot deg(u)}} \cdot (W^{(t)} \cdot h_u^{(t-1)})$
- 这里消息是经过归一化(使用了度矩阵)的邻居节点特征变换后的结果。
-
更新函数
- $h_v^{(t)} = \text{ReLU}(m_v^{(t)})$
- 非常简单,直接将聚合后的消息通过一个激活函数。
MPNN 的三个关键设计原则
- 置换不变性 (Permutation Invariance):消息聚合函数和读出函数必须对邻居的顺序不敏感(求和、均值、最大值),因为图没有天然的顺序。
- 局部性 (Locality):消息传递只发生在邻居节点之间,保证了模型的局部连接性,这与图的稀疏特性非常匹配。
- 可组合性 (Composability):通过堆叠多层消息传递,节点可以逐步获取更大范围的信息(感受野)。
MPNN 的优点与局限
优点
- 通用性强:几乎所有的 GNN 模型(GCN, GAT, GraphSAGE, GIN 等)都可以用 MPNN 框架来描述,只是具体的消息函数和更新函数不同。
- 表达能力:理论上,MPNN 可以逼近任意图上的函数,但其表达能力受限于消息聚合函数(如求和可能不如多重集合的区分能力强,GIN 论文对此有详细分析)。
- 适合分子等物理结构:在量子化学(预测分子能量、性质)等领域,MPNN 表现出色,因为它天然符合物理规律——原子的性质由它周围原子的相互作用决定。
局限
- 过平滑问题 (Over-smoothing):随着层数增加,所有节点的表示会趋向一致,失去区分度,这是所有 GNN 的通病。
- 长程依赖问题:要捕捉远距离节点的关系需要堆叠很多层,但这又会导致过平滑,注意力机制(GAT)或位置编码可以部分缓解,但仍然是活跃的研究领域。
- 计算复杂度:如果图比较稠密(邻居很多),消息聚合的计算量会很大,不过大多数实际图都是稀疏的。
代码示例:一个简单的 MPNN 层 (PyTorch 风格)
这里给出一个基础 MPNN 层的 TensorFlow/PyTorch 伪代码实现思路,让你理解它是如何工作的:
import torch
import torch.nn as nn
import torch.nn.functional as F
class SimpleMPNNLayer(nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
# 消息网络:为每个邻居学习一个变换
self.message_net = nn.Linear(in_features, out_features)
# 更新网络:将节点自身特征和聚合消息结合起来
self.update_net = nn.Linear(in_features + out_features, out_features)
def forward(self, h, adj_matrix):
"""
h: 节点特征矩阵 (N, in_features)
adj_matrix: 邻接矩阵 (N, N),描述连接关系,通常为稀疏矩阵以提高效率,这里简化用稠密示例
"""
# 1. 对每个节点,计算发送给邻居的消息(这里将节点特征变换一次)
# 实际上消息可以更复杂,(h_v, h_u, e_vu) 的函数
messages = self.message_net(h) # (N, out_features)
# 2. 消息聚合:通过邻接矩阵矩阵乘法实现邻居消息的求和
# 注意:实际应用中通常使用稀疏矩阵的 scatter_add
aggr_messages = torch.matmul(adj_matrix, messages) # (N, out_features)
# 3. 更新节点特征:将原始特征与聚合消息拼接后通过更新网络
combined = torch.cat([h, aggr_messages], dim=-1) # (N, in_features + out_features)
h_new = F.relu(self.update_net(combined)) # (N, out_features)
return h_new
-
如果你做分子性质预测,节点特征可以是原子的种类、电荷等;边特征可以是化学键类型;消息传递几层后,再通过读出层(例如对所有原子特征求和)得到分子图表示,最后用线性层输出预测值。
-
如果你做社交网络分析,节点特征是用户画像,边是关注关系,MPNN 可以学习每个用户的社交影响力表示。
MPNN 是一个优雅而强大的框架,它把图学习问题归结为“设计好的消息函数、更新函数和读出函数”,如果你理解了 MPNN,你就理解了 GNN 的核心逻辑。