3.4 MPNN:统一办案规程


3.4 MPNN:统一办案规程

本节摘要:消息传递神经网络(MPNN)不是又一件新装备,而是一册统一规程——把任意图模型拆解为消息、聚合、更新、读出诸相位,边特征升格为一等公民。本节按五段式展开:量子化学建模的统一动机、四相机制、相位公式、自定义消息传递层的实战代码、以及门控图网络等框架内变体。

案件背景:分子建模界的方言大战

委托来自计算化学界:预测分子的性质需要图神经网络,但各家实验室的模型——门控图网络、交互网络、分子神经指纹、谱方法变体——符号不同、接口不通,复现与改进都要先破译对方的"方言"。MPNN 的立项动机正是统一:作者证明这些模型都可以写成同一个消息传递框架的实例,差别只在各相位里填什么函数。这桩"标准化"案子对侦探社的价值同样直接——前三节的三件装备看起来各成体系,经 MPNN 一整理,你会发现它们只是同一规程的不同填空答案:GCN 的消息函数是线性变换、聚合是归一化求和;GraphSAGE 的聚合是抽样均值或池化;GAT 的消息带注意力系数。学完本节,读任何新图模型论文都变成"找它四个相位各填了什么"的填空题

机制拆解:四个相位

MPNN 把一轮完整的前向分成四个相位。消息相位:每条边上的源节点向目标节点发送消息,消息函数的输入是两端节点表示与这条边的特征——注意边特征在这里正式入场,键型、交易金额、关系类型都作为消息的原料,这在此前三个模型的原始版本里都是缺席的。聚合相位:目标节点把收到的消息按某种置换不变算子汇总(求和最常用)。更新相位:节点融合旧表示与聚合结果,生成新表示,常用门控循环单元以获得更精细的记忆管理。读出相位:跑完若干轮后,把全部节点表示汇总为整图的向量,读出函数必须对节点顺序不敏感(求和、注意力池化皆可),供图级任务使用。节点级任务可以跳过读出,直接取节点表示接任务头。

公式推导:相位式与既有模型的代入

四个相位各有标准记号。消息相位:m 等于消息函数作用于目标表示、源表示与边特征的拼接或组合。聚合相位:M 等于邻域内全部消息的求和(可替换为其他置换不变算子)。更新相位:新表示等于更新函数作用于旧表示、聚合消息(必要时加时间步信息)的组合。读出相位:图向量等于读出函数作用于全部节点表示的集合。

把已有模型代入这套记号,统一性立刻显现。GCN 的消息函数是"源表示乘权重矩阵"、聚合是"对称归一化系数加权的求和"、更新是"过非线性";GraphSAGE 的消息保持原样、聚合换成抽样均值或池化、更新是拼接投影;GAT 的消息乘上了注意力系数;门控图网络(GGNN)把更新相位填成门控循环单元并在消息里加入方向类型编码。模型间的全部差异被压缩进四个函数槽位——这就是 MPNN 作为"统一办案规程"的含义:它本身不提供新性能,提供的是设计与阅读的坐标系。

代码实战:自定义消息传递层

PyG 把 MPNN 的四相位做成了基类,自定义模型只需填空。下面写一个带边特征的定制层:消息函数把源节点特征与边特征相加。

import torch from torch_geometric.nn import MessagePassing class EdgeMessageConv(MessagePassing): """消息=源特征+边特征;聚合=求和;更新=直通加自环残差""" def __init__(self, in_dim, out_dim): super().__init__(aggr="add", flow="source_to_target") self.lin = torch.nn.Linear(in_dim, out_dim) def forward(self, x, edge_index, edge_attr): # propagate 内部自动完成:消息生成→按边分发→聚合 return self.propagate(edge_index, x=x, edge_attr=edge_attr) def message(self, x_j, edge_attr): # x_j 是 PyG 约定:按边展开的源节点特征(j 端) return self.lin(x_j + edge_attr) # 边特征直接进消息 def update(self, aggr_out, x): return aggr_out + x # 残差:旧档案不丢 d = 8 conv = EdgeMessageConv(d, d) x = torch.randn(6, d) edge_index = torch.tensor([[0,1,2,3,4],[1,2,3,4,5]]) edge_attr = torch.randn(5, d) # 每条边一份特征 out = conv(x, edge_index, edge_attr) print("输出形状:", out.shape) # 与输入同形,可堆叠多层

再用该层组装完整模型,体会"相位填空"的模块化:两层消息传递加读出,即成一个图分类器骨架。

import torch import torch.nn.functional as F from torch_geometric.nn import global_add_pool class MoleculeNet(torch.nn.Module): """MPNN 骨架:边特征版消息传递+求和读出(分子性质预测风格)""" def __init__(self, atom_dim, bond_dim, hid=32, out_dim=1): super().__init__() self.atom_emb = torch.nn.Linear(atom_dim, hid) self.bond_emb = torch.nn.Linear(bond_dim, hid) # 键特征对齐维度 self.conv1 = EdgeMessageConv(hid, hid) self.conv2 = EdgeMessageConv(hid, hid) self.head = torch.nn.Linear(hid, out_dim) def forward(self, x, edge_index, edge_attr, batch): h = self.atom_emb(x).relu() e = self.bond_emb(edge_attr) h = self.conv1(h, edge_index, e).relu() h = self.conv2(h, edge_index, e) g = global_add_pool(h, batch) # 读出相位:求和池化 return self.head(g) # 图级输出(如溶解度预测) atom_dim, bond_dim = 9, 3 # 原子独热+键型独热示意 model = MoleculeNet(atom_dim, bond_dim) x = torch.randn(12, atom_dim) ei = torch.tensor([[0,1,2,3,4,5,6,7,8,9,10],[1,2,3,4,5,6,7,8,9,10,11]]) ea = torch.randn(11, bond_dim) batch = torch.zeros(12, dtype=torch.long) # 全部原子属于同一个分子 print("分子预测输出:", model(x, ei, ea, batch).shape)

两个代码块合起来覆盖了四相位的完整实现路径:第一个填消息与更新槽位,第二个示范读出与任务头的接驳。实际分子项目中,原子特征表与键特征表从化学软件的解析结果构造,其余不变。

变体对比:框架内的填空答案

变体 消息槽位 更新槽位 特色与场景
GGNN 门控图网络 方向编码加线性 门控循环单元 收敛步数固定,程序验证、图到序列任务
交互网络 物理先验的消息函数 物理状态更新 物理仿真、粒子动力学
EdgeConv 边卷积 端点特征差过多层感知机 最大值聚合 点云分类(图视角看点云)
关系图卷积 按关系类型分消息函数 分关系聚合后求和 知识图谱(第六章重访)

统一规程在手,剩下一个悬而未决的问题:这些填空答案之间,有没有高下之分的理论标尺?下一节的鉴定科用"多重集合可区分性"给出裁决——GIN 及其背后的同构检验思想将回答"哪种聚合才是天花板"。


作者与出处
原作者: 灏天文库
来源:灏天文库
整理: 灏天文库整理
由灏天文库平台收录,内容或由平台用户上传,仅供学习交流
发布者: 作者: 灏天文库 转发
评论区 (0)
U