图神经网络 图神经网络(graph neural network,GNN)通过在相连节点之间传递消息来从图结构数据中学习。本文件涵盖消息传递框架、GCN、GraphSAGE、GIN、过平滑、图池化,以及节点级/边级/图级任务——这些核心架构支撑着分子性质预测、社交网络分析和推荐系统。 在前面几个文件里,我们奠定了数学基础:几何深度学习(文件 01)告诉我们要利用对称性,图论(文件 02)给我们提供了节点、边和邻接的语言。现在我们来构建直接在图上运算的神经网络。 核心挑战在于:图数据是不规则的。不像图像(固定网格)或序列(固定顺序),图的节点数目可变、连接方式可变、而且没有规范的节点顺序。一个为图设计的神经网络必须处理所有这些,同时还要对排列等变(重新标记节点不应改变输出)。
图神经网络(graph neural network,GNN)通过在相连节点之间传递消息来从图结构数据中学习。本文件涵盖消息传递框架、GCN、GraphSAGE、GIN、过平滑、图池化,以及节点级/边级/图级任务——这些核心架构支撑着分子性质预测、社交网络分析和推荐系统。
在前面几个文件里,我们奠定了数学基础:几何深度学习(文件 01)告诉我们要利用对称性,图论(文件 02)给我们提供了节点、边和邻接的语言。现在我们来构建直接在图上运算的神经网络。
核心挑战在于:图数据是不规则的。不像图像(固定网格)或序列(固定顺序),图的节点数目可变、连接方式可变、而且没有规范的节点顺序。一个为图设计的神经网络必须处理所有这些,同时还要对排列等变(重新标记节点不应改变输出)。
几乎所有 GNN 都遵循同一个配方,叫做消息传递(message passing,又称邻域聚合 neighbourhood aggregation)。这个想法简单又优雅:每个节点通过收集邻居的信息来更新自己的表示。
在每一层 l,每个节点 i 做三件事:
形式化地:
聚合 \bigoplus 必须是排列不变的(按什么顺序处理邻居无所谓),才能保证整体函数是排列等变的。这直接实现了文件 01 中的对称性原理。
经过 k 层消息传递之后,每个节点的表示编码了来自其 k 跳邻域的信息:所有在 k 条边以内可达的节点。第 1 层看到直接邻居,第 2 层看到邻居的邻居,以此类推。局部信息就是这样传播开来,构建出全局理解的。
GNN 的感受野随深度增长,就像 CNN 的感受野随层数增长一样(第 8 章)。但不同于在规则网格上的 CNN,GNN 每个节点的感受野形状会因图拓扑而异。
GCN(Kipf 和 Welling,2017)是奠基性的 GNN 架构。它把谱图卷积(来自文件 02)简化成一个优雅、高效的公式。
从谱卷积 g_\theta \star \mathbf{x} = U \, \text{diag}(\hat{g}_\theta) \, U^T \mathbf{x} 出发,Kipf 和 Welling 用一阶 Chebyshev 多项式来近似谱滤波器,从而完全避免了计算特征分解。化简之后,逐层的更新公式变成:
其中:
矩阵乘法 \hat{A} H^{(l)} 就是聚合步骤:对每个节点,它计算邻居特征(加上自环带来的自身特征)的加权平均。权重矩阵 W^{(l)} 是在所有节点之间共享的可学习变换。激活函数则加入非线性。
它出奇地简单:不过是一次矩阵乘法,再加上一个学习到的线性映射和激活函数。整个 GCN 层可以用一行代码写完。用 \tilde{D}^{-1/2} 做归一化是为了避免邻居太多的节点喧宾夺主:高度数的节点,其发出的消息会被相应地缩小。
在消息传递框架里,GCN 用的方案是:
GCN 是**直推式(transductive)的:它训练时需要整张图,无法处理新出现的、未见过的节点。如果一个新用户加入了社交网络,GCN 必须在整张图上重新训练。GraphSAGE(Hamilton 等,2017)用一种归纳式(inductive)**的方法解决了这个问题。
关键想法是邻域采样(neighbourhood sampling):不用全部邻居,而是采样一个固定大小的子集。这让计算与整张图的结构无关,从而能泛化到没见过的节点和图。
GraphSAGE 对节点 i 的更新是:
其中 \mathcal{S}(i) 是邻居的一个采样子集(例如从 500 个邻居里随机采样 10 个)。CONCAT 操作把节点自身特征与聚合后的邻居特征明确分开,让网络可以为"自己"和"邻域"学到不同的变换。
GraphSAGE 支持多种聚合函数:
这种采样策略让 GraphSAGE 可以扩展到非常大的图。训练时用节点的小批量:对每个目标节点,在第 1 层采样 k_1 个邻居,然后对其中每个邻居在第 2 层采样 k_2 个邻居。如果 k_1 = k_2 = 10、共 2 层,那么每个节点的计算树最多包含 10 \times 10 = 100 个节点,与整张图有多大无关。
不同的 GNN 架构有不同的表达能力(expressive power):即区分结构上不同的图的能力。GCN 和 GraphSAGE 尽管在实践中很有效,但在能区分哪些图结构上存在可证明的局限。
衡量 GNN 表达能力的理论工具是 Weisfeiler-Lehman(WL)检验,这是一个用于测试图同构(两张图在结构上是否完全相同)的经典算法。WL 检验通过把每个节点的标签与其邻居标签的多重集一起哈希来迭代地细化节点标签。
GIN(Xu 等,2019)被设计得与 WL 检验一样有表达能力,使它成为(在消息传递的理论极限之内)最强大的消息传递 GNN。关键洞见是:聚合函数对多重集必须是**单射(injective)**的(邻居特征的不同多重集必须产生不同的聚合值)。
求和聚合对多重集是单射的(把 \{1, 1, 2\} 求和得到 4,而 \{1, 3\} 也得到 4,但当特征向量维数足够多时,不同多重集的求和一般是不同的)。求平均和取最大值不是单射的:求平均无法区分 \{1, 1\} 和 \{2, 2\},取最大值无法区分 \{1, 2, 3\} 和 \{1, 1, 3\}。
GIN 的更新公式是:
它的机制很直观。每个消息传递层都把一个节点的特征与其邻居的特征做平均。经过许多轮平均之后,每个节点都已经看到(并混合了)它所在连通分量里的所有其他节点。特征变成了一个均匀的平均值,就像把一幅图像模糊太多次直到它变成一块纯色一样。
形式化地,反复施加归一化邻接 \hat{A} 会收敛到一个秩 1 矩阵(每一行都正比于图上随机游走的平稳分布)。这和幂迭代向主特征向量收敛是同一回事(第 2 章)。
过平滑把 GNN 限制在很浅的深度(通常是 2-4 层),不像 CNN 和 Transformer 可以从几十甚至几百层中获益。这意味着每个节点只能看到有限的邻域,这对需要长程信息的任务来说是个问题。
缓解办法包括:
对于图级任务(预测整张图的某个性质,比如分子的毒性),我们需要把所有节点的表示塌缩成一个图级向量。这就是图池化(graph pooling),是 CNN 中全局平均池化(第 8 章)的图类比。
最简单的办法是读出(readout):对所有节点特征组成的集合施加一个排列不变的函数:
这就是文件 01 里的 DeepSets 聚合,应用在最后一层 GNN 之后。求和保留了大小信息(100 个节点的图比 10 个节点的图求和结果更大),而求平均则把大小归一化掉。
**层级池化(hierarchical pooling)**逐步把图粗化,与 CNN 逐步对图像做下采样如出一辙。在每一层,一组组节点被合并成"超节点":
DiffPool(可微池化,Differentiable Pooling)学习一个软分配矩阵 S^{(l)} \in \mathbb{R}^{n_l \times n_{l+1}},把每个节点分配到一个簇:
这个分配矩阵由一个独立的 GNN 预测,使整个聚类过程端到端可微。这构造出一个层级:原始图 → 节点更少的粗化图 → 更粗的图 → 单个节点(即图的表示)。
TopKPool 采用更简单的办法:为每个节点学一个标量分数,保留分数最高的前 k 个节点,丢掉其余的。这是一种硬选择(不是软分配),计算上比 DiffPool 更便宜。
到目前为止的所有 GNN 都假设是同构图(homogeneous graph):只有一种节点类型、一种边类型。但大多数现实中的图是**异构图(heterogeneous graph)**的:有多种节点类型和多种边类型。一个知识图谱有人物节点、机构节点和地点节点,由"就职于""出生于""位于"等边连接。一个推荐系统有用户节点和商品节点,由"购买过""浏览过""评过分"等边连接。
异构图有一个模式(schema,又称元图 metagraph),定义了允许的节点类型和边类型。每种边类型连接一个特定的源类型到一个特定的目标类型。例如,"就职于"把 Person → Organisation 连起来。
关系 GCN(Relational GCN,R-GCN)(Schlichtkrull 等,2018)通过为每种边类型用一个独立的权重矩阵来处理异构边:
其中 \mathcal{R} 是边类型的集合,\mathcal{N}_r(i) 是通过关系 r 与节点 i 相连的邻居集合,W_r 是关系 r 专属的权重矩阵。自连接 W_0 单独处理节点自身的特征。
问题在于:关系类型一多,参数量就会爆炸(每种关系一个 d \times d 矩阵)。R-GCN 用**基分解(basis decomposition)**来缓解:W_r = \sum_{b=1}^{B} a_{rb} V_b,其中 V_b 是共享的基矩阵,a_{rb} 是每种关系的标量系数。这类似于低秩分解(第 2 章):关系专属的矩阵生活在一个低维子空间里。
异构图 Transformer(Heterogeneous Graph Transformer,HGT)(Hu 等,2020)把注意力机制用到异构图上。关键洞见是:注意力应当同时依赖于节点类型和连接它们的边类型。HGT 为 query、key、value 各自使用类型专属的投影矩阵:
其中 \tau(i) 是节点 i 的类型,\phi(i,j) 是它们之间的边类型。这保证模型对不同类型的关系用不同的方式分配注意力:一篇论文在注意到它的作者时,应当用与注意到它的参考文献时不同的注意力权重。
基于元路径(metapath)的方法定义出穿过模式的有意义的路径(例如 Author → Paper → Author 表示合著关系),并沿这些路径聚合信息。HAN(Heterogeneous Attention Network,异构注意力网络)在两个层级上应用注意力:每条元路径内部(这条路径上的哪些邻居重要?)和元路径之间(哪些关系模式重要?)。
**链接预测(link prediction)**问的是:给定已有的边,哪些缺失的边可能存在?这是知识图谱补全(预测缺失的事实)、推荐(预测用户会喜欢哪些商品)和社交网络分析(预测未来的朋友关系)的核心任务。
**基于嵌入的方法(embedding-based methods)**为每个实体学一个向量,为每种关系学一个变换,然后根据实体和关系契合的程度给潜在的边打分:
TransE 把关系建模为嵌入空间中的平移:如果 (h, r, t) 是一个有效的三元组(头实体、关系、尾实体),那么 \mathbf{h} + \mathbf{r} \approx \mathbf{t}。打分函数是 f(h, r, t) = -\|\mathbf{h} + \mathbf{r} - \mathbf{t}\|。直观上,关系向量在嵌入空间里把头实体"移动"到尾实体。
RotatE 把关系建模为复数空间中的旋转:\mathbf{t} = \mathbf{h} \circ \mathbf{r},其中 \circ 是逐元素复数乘法,且 |\mathbf{r}_i| = 1(模为 1 的复数就是旋转)。它能建模对称、反对称、反转和组合等模式,这些都是 TransE 做不到的。
ComplEx 使用复值嵌入和 Hermite 内积,使它能够建模非对称关系(如果 A 是 B 的老板,那么 B 不是 A 的老板)。
基于 GNN 的链接预测先用消息传递算出节点嵌入,再用端点嵌入给边打分。这把 GNN 的结构推理能力和嵌入方法的关系建模能力结合起来。GNN 编码器能捕捉到单一嵌入方法所遗漏的多跳邻域结构。
GNN 解决三类任务:
节点级任务:为每个节点预测一个性质。例子:在社交网络里给用户分类(机器人还是真人)、预测相互作用网络中每个蛋白质的功能、半监督节点分类(给少数节点打标签,预测其余的)。输出是把节点嵌入 \mathbf{h}_i^{(L)} 喂给一个分类器。
边级任务:为每条边预测一个性质,或预测某条边是否存在。例子:链接预测(这两个用户会成为好友吗?)、知识图谱补全(这两个实体之间是否存在这种关系?)、药物-药物相互作用预测。输出通常用到两个端点节点的嵌入:\hat{y}_{ij} = f(\mathbf{h}_i, \mathbf{h}_j),其中 f 是点积、拼接后接 MLP 或其他组合方式。
图级任务:为整张图预测一个性质。例子:分子性质预测(这个分子有毒吗?)、图分类(这个社交网络是不是机器人网络?)、图生成(设计一个具有所需性质的分子)。输出用图池化得到 \mathbf{h}_G,再对其做分类或回归。
import jax import jax.numpy as jnp # 图:5 个节点,一条带分支的链 A = jnp.array([[0, 1, 0, 0, 0], [1, 0, 1, 0, 0], [0, 1, 0, 1, 1], [0, 0, 1, 0, 0], [0, 0, 1, 0, 0]], dtype=float) # 加自环 A_hat = A + jnp.eye(5) D_hat = jnp.diag(A_hat.sum(axis=1)) D_inv_sqrt = jnp.diag(1.0 / jnp.sqrt(A_hat.sum(axis=1))) A_norm = D_inv_sqrt @ A_hat @ D_inv_sqrt # 节点特征:独热的单位矩阵 H = jnp.eye(5) # 权重矩阵(随机初始化) rng = jax.random.PRNGKey(0) W = jax.random.normal(rng, (5, 3)) * 0.5 # GCN 层:H' = ReLU(A_norm @ H @ W) H_new = jax.nn.relu(A_norm @ H @ W) print("Original features (one-hot):") print(H) print("\nAfter GCN layer:") print(jnp.round(H_new, 3)) print("\nNotice: connected nodes now have similar representations")
import jax.numpy as jnp # 两个不同的邻域多重集,但平均值相同 # 节点 A:邻居特征是 [1, 1, 1, 1](四个邻居,全为 1) # 节点 B:邻居特征是 [2, 2] (两个邻居,全为 2) neighbours_A = jnp.array([[1.0], [1.0], [1.0], [1.0]]) neighbours_B = jnp.array([[2.0], [2.0]]) # 求平均聚合 mean_A = neighbours_A.mean(axis=0) mean_B = neighbours_B.mean(axis=0) print(f"Mean A: {mean_A}, Mean B: {mean_B}, Same: {jnp.allclose(mean_A, mean_B)}") # 求和聚合 sum_A = neighbours_A.sum(axis=0) sum_B = neighbours_B.sum(axis=0) print(f"Sum A: {sum_A}, Sum B: {sum_B}, Same: {jnp.allclose(sum_A, sum_B)}") print("\nSum distinguishes these multisets; mean does not!")
import jax.numpy as jnp import matplotlib.pyplot as plt # 随机图 A = jnp.array([[0,1,1,0,0,0], [1,0,1,0,0,0], [1,1,0,1,0,0], [0,0,1,0,1,1], [0,0,0,1,0,1], [0,0,0,1,1,0]], dtype=float) A_hat = A + jnp.eye(6) D_inv_sqrt = jnp.diag(1.0 / jnp.sqrt(A_hat.sum(axis=1))) A_norm = D_inv_sqrt @ A_hat @ D_inv_sqrt # 初始特征:每个节点各不相同 H = jnp.array([[1,0], [0,1], [1,1], [-1,0], [0,-1], [-1,-1]], dtype=float) distances = [] for k in range(20): H = A_norm @ H # 衡量特征有多分明(跨节点的标准差) spread = jnp.std(H, axis=0).mean() distances.append(float(spread)) plt.plot(distances, "o-") plt.xlabel("Number of message-passing rounds") plt.ylabel("Feature spread (std across nodes)") plt.title("Over-Smoothing: Features Converge with Depth") plt.show()