\[ \begin{align}\begin{aligned}\newcommand{\ba}{\boldsymbol{a}} \newcommand{\bb}{\boldsymbol{b}} \newcommand{\be}{\boldsymbol{e}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bw}{\boldsymbol{w}} \newcommand{\bx}{\boldsymbol{x}} \newcommand{\by}{\boldsymbol{y}} \newcommand{\bz}{\boldsymbol{z}} \newcommand{\bd}{\boldsymbol{d}} \newcommand{\bv}{\boldsymbol{v}} \newcommand{\bs}{\boldsymbol{s}}\\\newcommand{\btheta}{\boldsymbol{\theta}} \newcommand{\bbeta}{\boldsymbol{\beta}} \newcommand{\bgamma}{\boldsymbol{\gamma}} \newcommand{\bsigma}{\boldsymbol{\sigma}} \newcommand{\md}{\mbox{d}} \newcommand{\bmu}{\boldsymbol{\mu}} \newcommand{\bone}{\boldsymbol{1}} \newcommand{\bzero}{\boldsymbol{0}} \newcommand{\bepsilon}{\boldsymbol{\epsilon}} \newcommand{\bphi}{\boldsymbol{\phi}} \newcommand{\bh}{\boldsymbol{h}} \newcommand{\bc}{\boldsymbol{c}} \newcommand{\br}{\boldsymbol{r}} \newcommand{\bQ}{\boldsymbol{Q}} \newcommand{\bK}{\boldsymbol{K}} \newcommand{\bV}{\boldsymbol{V}} \newcommand{\bSigma}{\boldsymbol{\Sigma}} \newcommand{\bg}{\boldsymbol{g}} \newcommand{\bxi}{\boldsymbol{\xi}} \newcommand{\bvarepsilon}{\boldsymbol{\varepsilon}} \newcommand{\bdelta}{\boldsymbol{\delta}} \newcommand{\bq}{\boldsymbol{q}} \newcommand{\bk}{\boldsymbol{k}} \newcommand{\bJ}{\boldsymbol{J}} \newcommand{\bp}{\boldsymbol{p}} \newcommand{\bi}{\boldsymbol{i}} \newcommand{\bo}{\boldsymbol{o}} \newcommand{\bE}{\boldsymbol{E}} \newcommand{\bH}{\boldsymbol{H}} \newcommand{\bL}{\boldsymbol{L}} \newcommand{\bu}{\boldsymbol{u}} \newcommand{\bLambda}{\boldsymbol{\Lambda}} \newcommand{\trans}{^{\rm\scriptsize T}} \newcommand{\var}{\mathrm{var}}\\\newcommand{\bA}{\boldsymbol{A}} \newcommand{\bB}{\boldsymbol{B}} \newcommand{\bC}{\boldsymbol{C}} \newcommand{\bD}{\boldsymbol{D}} \newcommand{\bG}{\boldsymbol{G}} \newcommand{\bI}{\boldsymbol{I}} \newcommand{\bM}{\boldsymbol{M}} \newcommand{\bP}{\boldsymbol{P}} \newcommand{\bS}{\boldsymbol{S}} \newcommand{\bU}{\boldsymbol{U}} \newcommand{\bW}{\boldsymbol{W}} \newcommand{\bX}{\boldsymbol{X}} \newcommand{\bY}{\boldsymbol{Y}} \newcommand{\bZ}{\boldsymbol{Z}} \newcommand{\cotp}{\textcolor[RGB]{48,209,88}{TP}} \newcommand{\cotn}{\textcolor[RGB]{100,210,255}{TN}} \newcommand{\cofp}{\textcolor[RGB]{94,92,230}{FP}} \newcommand{\cofn}{\textcolor[RGB]{191,90,242}{FN}}\\\newcommand{\numcotp}{\textcolor[RGB]{48,209,88}{50}} \newcommand{\numcotn}{\textcolor[RGB]{100,210,255}{30}} \newcommand{\numcofp}{\textcolor[RGB]{94,92,230}{10}} \newcommand{\numcofn}{\textcolor[RGB]{191,90,242}{10}} \DeclareMathOperator*{\argmin}{arg\,min}\end{aligned}\end{align} \]

消息传递神经网络#

学习目标与记号#

  1. 用消息、聚合和更新三个步骤统一表达常见图神经网络层;

  2. 证明置换不变聚合器带来的节点置换等变性,并分析多层感受野;

  3. 实现带自环和零度保护的邻居聚合,并比较循环、稀疏和批量实现;

  本节沿用统一记号:普通小写字母表示标量,粗体小写字母表示向量,粗体大写字母表示矩阵或高阶张量;层编号写作上标 \([l]\)⁠,节点、样本或边编号写作下标;转置写作 \(\trans\)⁠。除非另有说明,节点特征按行存放,邻接约定为 \(A_{ij}>0\) 表示节点 \(j\) 向节点 \(i\) 发送消息。正文与练习中的程序都应同时检查数值结果和数组维度。本节练习的参考答案见 消息传递神经网络答案⁠。

  开始推导前,应先明确 图矩阵与特征记号 以及 置换等变性⁠;后续 GCN⁠、GATGraphSAGE 都是消息传递框架的具体化。

  消息传递神经网络(message passing neural network,MPNN)把大量图模型写成同一个局部计算框架。第 \(l\) 层中,节点 \(v\) 先从邻居 \(u\in\mathcal{N}(v)\) 接收消息,再以不依赖邻居排列的方式聚合,最后与自身旧状态一起更新。

统一公式#

  令 \(\bh_v^{[0]}=\bx_v\)⁠,则一层消息传递可写为

\[\begin{split}\begin{aligned} \boldsymbol{m}_{u\to v}^{[l]} &=\psi^{[l]}\!\left( \bh_v^{[l-1]}, \bh_u^{[l-1]}, \be_{uv}\right),\\ \boldsymbol{m}_v^{[l]} &=\operatorname{AGG}^{[l]} \left(\left\{\boldsymbol{m}_{u\to v}^{[l]}: u\in\mathcal{N}(v)\right\}\right),\\ \bh_v^{[l]} &=\phi^{[l]}\!\left( \bh_v^{[l-1]},\boldsymbol{m}_v^{[l]}\right). \end{aligned}\end{split}\]

消息函数 \(\psi\) 可以使用发送节点、接收节点和边特征,其中上式中的 \(\be_{uv}\) 即边 \((u,v)\) 的特征向量;聚合器 \(\operatorname{AGG}\) 常用求和、均值、最大值或注意力加权和;更新函数 \(\phi\) 可以是线性层、MLP、门控单元或带残差的模块。GCN 固定使用度归一化权重,GAT 学习边上的注意力权重,GraphSAGE 强调采样和归纳聚合。

  下面的动画以一个目标节点为例,依次展示邻居消息的计算、汇总以及节点表示的更新过程,并说明堆叠多层后节点能够利用更远处的信息。

聚合器与信息保留#

  求和能保留邻居数量信息,但数值尺度会随度数增长;均值控制尺度,却可能把“一个值为 1 的邻居”和“两个值为 1 的邻居”映射到相同结果;最大值对重复元素不敏感。不存在对所有任务都最优的聚合器,选择应结合任务需要区分的邻域结构。

  当消息只是发送节点状态时,求和聚合的矩阵形式为

\[\bM^{[l]}=\bA\cdot\bH^{[l-1]},\]

其中,\(\bM^{[l]}\in\mathbb{R}^{n\times d^{[l-1]}}\) 是这一层所有节点消息按行堆叠成的矩阵,其第 \(i\) 行正是节点 \(i\) 收到的消息;若第 \(i\) 个节点的第 \(l-1\) 层表示为列向量 \(\bh_i^{[l-1]}\)⁠,则 \(\bH^{[l-1]}=[(\bh_1^{[l-1]})\trans;\ldots;(\bh_n^{[l-1]})\trans]\in\mathbb{R}^{n\times d^{[l-1]}}\)⁠。按照本章的邻接约定,第 \(i\) 行得到所有 \(A_{ij}>0\) 的节点 \(j\) 的加权和。使用有向图时,必须明确聚合入边、出边还是双向边。

感受野、深度与读出#

消息传递神经网络的消息、聚合、更新和多跳感受野

图 36 一层聚合一跳邻居;堆叠多层后,节点表示可以依赖更远的节点#

  一层只使用一跳邻域。由归纳法可知,堆叠 \(L\) 层后,节点 \(v\) 的表示最多依赖 \(L\) 跳可达节点。更深并不总是更好:表示可能逐渐相似而过度平滑,远距离大量信息也可能被压缩进固定维度向量而出现过度挤压。残差、跳层连接、位置或结构编码,以及改变传播范围,都可用于缓解这些问题。

  节点级任务可以直接使用各节点的最终表示进行预测,而 图级任务 需要先把整张图中数量可能不同的节点表示汇总起来。完成这一汇总的函数称为读出函数(readout function)。它以最后一层的全部节点表示为输入,并输出一个维度固定的整图表示:

\[\bh_{\mathcal{G}} =\operatorname{READOUT} \left(\{\bh_v^{[L]}:v\in\mathcal{V}\}\right),\]

其中,\(\{\bh_v^{[L]}:v\in\mathcal{V}\}\) 表示图 \(\mathcal{G}\) 中所有节点在第 \(L\) 层得到的表示,\(\bh_{\mathcal{G}}\) 是用于概括整张图信息的固定维度向量。不同图的节点数可以不同,但只要节点表示的维度相同,读出后得到的 \(\bh_{\mathcal{G}}\) 就具有相同维度,因而可以继续交给同一个预测函数处理。由于节点编号不代表实际含义,交换节点的排列顺序不应改变读出结果;因此图级任务使用的读出函数必须具有置换不变性。

  最常见的三种读出形式为

\[\begin{split}\begin{aligned} \bh_{\mathcal{G},\mathrm{sum}} &=\sum_{v\in\mathcal{V}}\bh_v^{[L]},\\ \bh_{\mathcal{G},\mathrm{mean}} &=\frac{1}{|\mathcal{V}|}\sum_{v\in\mathcal{V}} \bh_v^{[L]},\\ \left(\bh_{\mathcal{G},\mathrm{max}}\right)_j &=\max_{v\in\mathcal{V}} \left(\bh_v^{[L]}\right)_j. \end{aligned}\end{split}\]

  求和读出把所有节点表示相加,对节点数量和重复出现的特征更敏感,但结果的数值尺度容易随节点数增加;均值读出先求和再除以节点数,便于比较大小不同的图,但可能丢失节点数量信息;最大值读出在每个维度上选择所有节点中的最大值,适合判断某类强特征是否出现,却不能说明该特征出现了多少次。更复杂的注意力读出会根据节点表示学习权重,再计算加权和;只要权重计算和加权求和不依赖节点的排列顺序,读出结果仍具有置换不变性。

  得到整图表示后,可以再用预测函数 \(g_{\btheta}\) 完成图级分类或回归:

\[\widehat{\by}_{\mathcal{G}} =g_{\btheta}\!\left(\bh_{\mathcal{G}}\right),\]

其中,\(g_{\btheta}\) 通常是线性层或多层全连接神经网络。对于分子图,可以先用消息传递计算每个原子的表示,再用读出函数得到整个分子的表示,最后预测分子的类别、溶解度或毒性。在一个批次同时处理多张图时,程序需要先记录每个节点属于哪张图,再分别对每张图中的节点进行读出,不能把不同图的节点混在一起汇总。

如何选择读出函数

  若任务与节点数量或某类结构出现的次数有关,可以优先比较求和读出;若希望减小图大小不同造成的数值尺度差异,可以比较均值读出;若任务主要关心某个显著特征是否出现,可以比较最大值读出。实际应用中应在其他条件相同的情况下比较不同读出函数,并根据验证集表现作出选择。

实现均值聚合#

  下面用 NumPy 实现一层均值聚合,把前面的矩阵公式转换为程序。输入 adjacency 是大小为 \(n\times n\) 的邻接矩阵,按照本章约定,第 \(i\) 行记录哪些节点向节点 \(i\) 发送信息;features 是大小为 \(n\times d\) 的节点特征矩阵;add_self_loops 决定是否把节点自身也加入参与求平均的节点集合。函数返回大小仍为 \(n\times d\) 的矩阵,其中第 \(i\) 行是节点 \(i\) 汇总自身及邻居信息后得到的新表示。

import numpy as np

def mean_aggregate(adjacency, features, add_self_loops=True):
    adjacency = np.asarray(adjacency, dtype=float)
    features = np.asarray(features, dtype=float)
    if adjacency.ndim != 2 or adjacency.shape[0] != adjacency.shape[1]:
        raise ValueError("邻接矩阵必须是方阵")
    if features.ndim != 2 or adjacency.shape[0] != features.shape[0]:
        raise ValueError("邻接矩阵与节点特征不匹配")
    a = adjacency.copy()
    if add_self_loops:
        a += np.eye(len(a))
    degree = a.sum(axis=1, keepdims=True)
    if np.any(degree == 0):
        raise ValueError("存在没有自环的零度节点")
    return (a @ features) / degree

  代码先检查邻接矩阵和节点特征矩阵的维度是否正确,再根据 add_self_loops 的取值决定是否给每个节点增加一条指向自己的边。矩阵乘法 a @ features 对每个节点收到的特征求和;degree 记录每个节点参与求和的总权重;最后逐行除以 degree便得到相应的平均值。对度数是否为 0 的检查可以避免除以 0;加入自环后,原本没有邻居的节点也能够保留自身特征。

  这个实现使用完整的邻接矩阵,适合在小图上手算或检查程序结果,不适合直接处理大图。训练大图时,通常只沿实际存在的边汇总信息,或者使用框架提供的稀疏矩阵乘法,从而避免保存含有大量 0 的完整邻接矩阵。

Shiny 交互演示:消息传递、GCN 与 GAT

  交互页面在同一张小图和同一组节点特征上比较均值消息传递、GCN 对称归一化聚合与单头 GAT。选择一个接收节点后,可以依次查看入边、聚合权重、每个邻居的贡献和更新后的节点表示。页面使用单位权重矩阵突出传播过程,不是经过训练的节点分类器。

点击打开“图神经网络逐步实验室”

核心推导与实现核验#

核心关系

\[\boldsymbol{m}_v^{[l]}=\operatorname{AGG}_{u\in\mathcal{N}(v)}\psi^{[l]}(\bh_v^{[l-1]},\bh_u^{[l-1]},\be_{uv}),\quad \bh_v^{[l]}=\phi^{[l]}(\bh_v^{[l-1]},\boldsymbol{m}_v^{[l]}).\]

  推导路径。 先沿每条入边计算消息,再在接收节点处按邻居集合归约,最后把聚合消息与节点旧状态交给共享更新函数。共享参数和对集合不敏感的聚合共同保证对节点编号的正确响应。

关键条件

  一层输出只依赖节点自身和一跳邻居。假设第 \(l-1\) 层只依赖 \(l-1\) 跳邻域,第 \(l\) 层再从一跳邻居读取这些表示,其依赖集合包含在 \(l\) 跳邻域内。因此,由归纳法可知,\(L\) 层最多依赖 \(L\) 跳邻域。

数据规模

  若图中有 \(n\) 个节点,上一层每个节点用 \(d^{[l-1]}\) 个数表示,消息包含 \(d_m\) 个数,新表示包含 \(d^{[l]}\) 个数,那么整批节点表示分别为 \(n\times d^{[l-1]}\)\(n\times d^{[l]}\)⁠。第一行到第 \(n\) 行始终要与同一批节点一一对应。

常见误区

  聚合器必须处理空邻域并明确自环;有向边方向写反会改变信息流;在循环中原地覆盖节点状态会让同一层不同节点使用不同时间的值。

动手检查

  用三节点链手算一层和两层结果;随机置换节点编号检查等变性;循环、稠密矩阵和稀疏边聚合应在相同约定下给出一致结果。

数值稳定性与规模

  均值聚合的分母要保护零度节点;求和聚合在高阶节点上可能尺度过大。深层训练应监控表示方差、节点间余弦相似度和梯度范数。

本节小结#

  1. MPNN 将图层分成消息、聚合和更新三个可独立检查的步骤。

  2. 置换不变聚合器配合共享参数产生节点置换等变性。

  3. 层数扩大感受野,也可能带来过度平滑、过度挤压和计算扩张。

综合练习#

  程序题应固定随机种子、写出维度断言并报告运行环境。全部参考答案见 消息传递神经网络答案⁠。

  1. 两层均值聚合计算。 对无向三节点链 \(1-2-3\)⁠,初始标量特征为 \((1,2,4)\)⁠。每层加入自环,只做均值聚合且不使用其他线性变换或非线性。计算第一层和第二层三个节点的表示,并说明第二层各节点最多使用了几跳信息。

  2. 置换等变性证明。 设消息函数和更新函数在所有节点、边上共享参数,聚合器对邻居枚举顺序不敏感。证明同步置换邻接关系和节点特征后,一层节点输出按同一置换变化;再说明使用 sum 或 mean 图级读出为何得到置换不变结果。

  3. 感受野证明。 用数学归纳法证明:不含全局读出或跨图捷径的 \(L\) 层消息传递网络中,节点 \(v\) 的表示最多依赖 \(L\) 跳可达节点。说明“最多依赖”为什么不等于网络一定能够区分所有 \(L\) 跳结构。

  4. 聚合器计算与信息损失。 分别计算多重集合 \(\{1\}\)⁠、\(\{1,1\}\)⁠、\(\{1,3\}\)⁠、\(\{2,2\}\)\(\{0,2\}\) 的 sum、mean 与 max。各给出一个碰撞例子,说明三种聚合器分别保留或丢失了什么信息。

  5. 稀疏均值聚合实现。 基于维度为 \(2\times m\)edge_index 实现带自环的 mean 聚合,不构造 \(n\times n\) 邻接矩阵。程序应按接收端点累加消息和度数,处理重复边、孤立节点和可选边权,并返回聚合结果与实际度数。

  6. 等价性与置换测试。 在若干小图上比较逐边循环、稠密矩阵和第 5 题稀疏实现,逐节点检查结果;再随机置换节点编号,验证节点输出等变、sum/mean 图读出不变。测试必须包含有向单边、孤立节点、重复边和非法端点。

  7. 可复用 MPNN 层。 实现可选择 sum、mean 或 max 的 MPNN 层,消息函数可读取发送节点、接收节点和边特征,更新函数可选择残差连接。编写测试确保同一层不原地覆盖旧状态、边方向符合约定、批量图之间不传递消息,并用 自动微分 检查所有可训练参数都获得有限梯度。

  8. MLP 与 MPNN。 构造或选择一个标签确实依赖邻域结构的节点二分类任务,比较只使用节点特征的 MLP 与单层、两层 MPNN。固定训练/验证/测试划分、随机种子集合、近似参数量、训练更新次数和每次更新的目标节点数;报告准确率或 macro-F1⁠、训练时间、固定批量预测时间、峰值内存、参数量和消息函数调用次数。

  9. 循环、稠密与稀疏实现。 对同一组稀疏图比较逐边循环、稠密矩阵乘法和 scatter/reduce 三种 mean 聚合。固定图、特征、随机种子集合、数据类型、预热次数和重复计时次数,并处理完全相同的边;报告数值误差、固定图批量的预测时间、峰值内存和实际处理边数,同时改变 \(n\) 与图密度观察扩展趋势。

  10. 深度与聚合器。 在同一节点分类任务上比较层数 \(L\in\{1,2,4,8\}\) 以及 sum、mean 聚合。各组使用相同的数据划分、随机种子、隐藏宽度、训练更新次数和评价程序;比较有无自环时只改变自环设置,比较有无残差连接时只改变残差连接设置。报告预测质量、训练时间、固定批量预测时间、峰值内存、消息调用次数、节点表示方差和节点间平均余弦相似度,并根据结果分别分析层数、聚合方法、自环和残差连接可能带来的影响。