图注意力网络 图注意力网络(graph attention network,GAT)用学习到的、依赖数据的加权来取代均匀的邻域聚合。本文件涵盖 GAT、多头图注意力、GATv2、图 Transformer(Graph Transformer)、位置编码与结构编码,以及可扩展性。 在 GCN(文件 03)里,每个节点用由图结构(归一化邻接)决定的固定权重来聚合邻居特征。一个有三个邻居的节点给每个邻居大致相等的权重($\approx 1/3$)。但并非所有邻居都同样重要:来自密切合作者的消息应当比来自点头之交的更有分量。 图注意力网络(Graph Attention Networks)通过学习该关注哪些邻居来解决这个问题,用的正是驱动 Transformer 的那个注意力机制(第 7 章)。
图注意力网络(graph attention network,GAT)用学习到的、依赖数据的加权来取代均匀的邻域聚合。本文件涵盖 GAT、多头图注意力、GATv2、图 Transformer(Graph Transformer)、位置编码与结构编码,以及可扩展性。
在 GCN(文件 03)里,每个节点用由图结构(归一化邻接)决定的固定权重来聚合邻居特征。一个有三个邻居的节点给每个邻居大致相等的权重(\approx 1/3)。但并非所有邻居都同样重要:来自密切合作者的消息应当比来自点头之交的更有分量。
图注意力网络(Graph Attention Networks)通过学习该关注哪些邻居来解决这个问题,用的正是驱动 Transformer 的那个注意力机制(第 7 章)。每个节点不再使用固定的、基于结构的权重,而是对其邻居计算动态的、基于内容的注意力分数。
其中 W \in \mathbb{R}^{d' \times d} 是一个共享的线性变换,\| 表示拼接,\mathbf{a} \in \mathbb{R}^{2d'} 是一个可学习的注意力向量。分数 e_{ij} 衡量节点 j 的特征对节点 i 有多重要。
原始分数用 softmax 在所有邻居上做归一化:
与 GCN 的关键区别在于:权重 \alpha_{ij} 是从数据中学到的,而不是由图结构固定的。一个节点可以学会聚焦于最有信息量的邻居,同时忽略噪声大或无关的邻居。
注意,注意力只在边上计算(节点 i 只注意到它的邻居 \mathcal{N}(i)),而不是在所有节点对上计算。这保证计算量与边数成正比,而不是与节点数的平方成正比。
每个头可以注意到邻域的不同方面:一个头可能聚焦于结构特征,另一个聚焦于语义相似性。这与 Transformer 中多头注意力的动机相同:不同的头捕捉不同类型的关系。
当有 K 个头、每个头输出维度为 d' 时,拼接后的输出维度是 K \times d'。最后一层通常用求平均而不是拼接,以产生一个固定大小的输出。
原始的 GAT 有一个微妙的局限:它的注意力函数是静态的(static,又称基于排序的 ranking-based)。注意力分数依赖于拼接 [W\mathbf{h}_i \| W\mathbf{h}_j],但由于注意力向量 \mathbf{a} 是在拼接之后才施加的,它可以被分解成两个独立的成分:\mathbf{a}^T [W\mathbf{h}_i \| W\mathbf{h}_j] = \mathbf{a}_1^T W\mathbf{h}_i + \mathbf{a}_2^T W\mathbf{h}_j。
这意味着对于给定的节点 i,其邻居的排序完全由邻居自身的特征 \mathbf{h}_j 决定(项 \mathbf{a}_1^T W\mathbf{h}_i 对 i 的所有邻居来说是常数)。注意力排序并没有真正依赖于发出查询的节点的特征。节点 i 和节点 k 会把同一组邻居排出完全相同的顺序,这限制了表达能力。
GATv2(Brody 等,2022)通过把非线性放在注意力向量之前来修复这一点:
标准的消息传递 GNN 受图拓扑的限制:一个节点只能注意到它的直接邻居。经过 k 层之后,来自 k 跳邻居的信息已经被多次聚合步骤混合,损失了保真度。这种局部瓶颈(加上文件 03 中的过平滑)限制了对长程依赖的捕捉能力。
图 Transformer(Graph Transformers)通过对所有节点对施加全局自注意力来打破这个瓶颈,无论它们之间是否有边。每个节点都能在单层内注意到所有其他节点,就像在标准 Transformer 里一样(第 7 章)。
基本想法:把所有节点当作 token,施加 Transformer 的自注意力:
其中 Q = XW_Q,K = XW_K,V = XW_V 是节点特征 X 的 query、key、value 投影(与第 7 章完全一样)。这相当于在一个完全图(完全图 K_n,文件 02)上的 GNN。
问题在于:完全图忽略了实际的图结构。边的信息(谁真正与谁相连)丢失了。有两种方法可以把它找回来:
Graphormer(Ying 等,2021)通过在注意力分数里加入偏置项来把图结构注入 Transformer:
空间偏置 b_{\text{spatial}} 编码节点 i 与 j 之间的最短路径距离。边偏置 b_{\text{edge}} 编码沿最短路径上的边特征。此外,Graphormer 还使用中心性编码(centrality encoding),把节点的度加到它的输入嵌入里,让模型知道每个节点的结构角色。
GPS(General, Powerful, Scalable Graph Transformer,Rampášek 等,2022)在每一层里把局部消息传递和全局注意力结合起来:
序列上的 Transformer 用位置编码(第 7 章)注入顺序信息。图没有规范的顺序,所以需要图专属的编码。
**拉普拉斯特征向量编码(Laplacian eigenvector encodings)**用图拉普拉斯(文件 02)的特征向量作为位置特征。k 个最小的非平凡特征向量给出图的一个谱嵌入:在图上"相近"的节点有相似的特征向量值。这些向量被拼接到节点特征上。
一个微妙之处:拉普拉斯特征向量有符号歧义(如果 \mathbf{u} 是特征向量,-\mathbf{u} 也是)。模型必须对这些符号翻转保持不变。解决办法包括在训练时把随机符号翻转当作数据增强,或者学习对符号不变的变换。
**随机游走编码(random walk encodings)**计算从节点 i 出发的随机游走经过 k 步后回到 i 的概率,其中 k = 1, 2, \ldots, K。这些概率编码了局部结构信息:处于密集簇中的节点返回概率高,而处于稀疏区域的节点返回概率低。落地概率 p_{ii}^{(k)} = (A_{\text{rw}}^k)_{ii},其中 A_{\text{rw}} = D^{-1}A 是随机游走转移矩阵。
**度编码(degree encodings)**干脆把节点度作为特征加进去。这出奇地有效,因为度是一个很强的结构信号:叶节点(度为 1)、桥节点和枢纽节点行为各不相同。
这些编码提供了普通 Transformer 所缺少的结构信息,使图 Transformer 在需要长程推理的任务上能胜过标准的消息传递 GNN。
GNN 的根本性可扩展性挑战在于,图可能有数百万节点和数十亿条边。在整张图上训练一个 GNN 需要把所有节点特征和整个邻接矩阵都放进内存,这往往不可行。
GNN 的**小批量训练(mini-batch training)比图像或序列更复杂,因为节点之间相互连接。简单地采样一批节点会牵扯出它们的邻居(第 1 层)、邻居的邻居(第 2 层),等等。这种邻域爆炸(neighbourhood explosion)**意味着 1000 个目标节点的一批,其计算图里可能牵涉到数百万个节点。
邻域采样(GraphSAGE 风格,文件 03)通过每层每节点只采样固定数目的邻居来限制爆炸。在 2 层、每层 15 个采样的情况下,每个目标节点的子图最多有 15^2 = 225 个节点,与整张图的大小无关。
Cluster-GCN(Chiang 等,2019)用图聚类算法(例如 METIS)把图划分成若干簇,然后一次在一个簇上训练。簇内的边很密(大多数邻居都在同一个簇里),所以子图能捕捉到相关结构。簇间的边则通过偶尔纳入跨簇边来处理。
图 Transformer 的可扩展性更难,因为全局注意力是 O(n^2)。对于有数百万节点的图,完整注意力不可行。解决办法包括:
我们到目前为止研究的图都是静态的(static):节点、边和特征是固定的。但许多现实中的图随时间演化:新用户加入社交网络、金融交易创造出边、交通模式在一天里不断变化、分子相互作用此起彼伏。
**时序图(temporal graph)**给每条边加一个时间戳:(i, j, t) 表示节点 i 在时刻 t 与节点 j 互动。挑战在于学到既能捕捉图结构、又能捕捉时间动态的表示。
这里有两种范式:
离散时间动态图(discrete-time dynamic graphs,DTDG):图被表示成一串快照 G_1, G_2, \ldots, G_T,每个时间步一张。一个 GNN 处理每张快照,再用一个 RNN 或时序注意力机制捕捉快照之间的演化。这很简单,但会丢失细粒度的时间信息(快照之间的事件会丢失),还需要选择快照频率。
连续时间动态图(continuous-time dynamic graphs,CTDG):事件被建模为一个带时间戳的互动流。每个事件 (i, j, t) 在它发生的精确时刻更新节点 i 和 j 的表示。这保留了所有时间信息。
时序图网络(Temporal Graph Network,TGN)(Rossi 等,2020)是领先的 CTDG 架构。每个节点维护一个记忆状态(memory state) \mathbf{s}_i(t),每当该节点参与一次互动时它就被更新:
其中 \mathbf{m}_i(t) 是从该次互动算出的消息(结合了两个节点的特征、边特征和时间编码)。GRU(第 6 章)选择性地保留和遗忘过去的信息,使记忆既能捕捉长期模式,又能适应最近的事件。
**时间编码(time encoding)**把自上次互动以来流逝的时间表示成一个特征向量,类似于 Transformer 中的位置编码(第 7 章)。一种常见做法使用可学习的傅里叶特征:
这给模型一个关于时间间隔的丰富表示:"这个用户 5 分钟前还活跃"和"3 个月前还活跃"会被嵌入成不同的向量。
**时序图注意力(Temporal Graph Attention,TGAT)**在一个节点的时序邻域(即最近的一组互动)上施加自注意力,每个互动既按特征相关性(像 GAT 那样)加权,又按时间上的临近程度加权。来自遥远过去的互动自然会被降权。
应用包括欺诈检测(金融图中异常的交易模式)、交通预测(从历史流量模式预测拥堵)、社交网络动态(预测病毒式内容的传播),以及随时间变化的药物相互作用预测。
import jax import jax.numpy as jnp rng = jax.random.PRNGKey(0) k1, k2, k3 = jax.random.split(rng, 3) n_nodes, d_in, d_out = 5, 4, 3 # 随机节点特征 H = jax.random.normal(k1, (n_nodes, d_in)) # 可学习参数 W = jax.random.normal(k2, (d_in, d_out)) * 0.5 a = jax.random.normal(k3, (2 * d_out,)) * 0.5 # 邻接(节点 0 连到 1、2、3) neighbours_of_0 = [1, 2, 3] # 变换特征 Wh = H @ W # (n_nodes, d_out) # 计算节点 0 的注意力分数 h_i = Wh[0] scores = [] for j in neighbours_of_0: h_j = Wh[j] e_ij = jnp.dot(a, jnp.concatenate([h_i, h_j])) e_ij = jax.nn.leaky_relu(e_ij, negative_slope=0.2) scores.append(float(e_ij)) scores = jnp.array(scores) alpha = jax.nn.softmax(scores) print(f"Raw scores: {scores}") print(f"Attention weights: {alpha}") print(f"Sum of weights: {alpha.sum():.4f}") # 加权聚合 h_new = sum(alpha[k] * Wh[neighbours_of_0[k]] for k in range(len(neighbours_of_0))) print(f"Updated node 0 features: {h_new}")
import jax import jax.numpy as jnp # 4 个节点:节点 0 连到 1、2、3 A = jnp.array([[0,1,1,1], [1,0,0,0], [1,0,0,0], [1,0,0,0]], dtype=float) # 特征:节点 1 很相关,节点 2 是噪声,节点 3 中等 H = jnp.array([[0.0, 0.0], # 节点 0 [1.0, 0.0], # 节点 1(信号) [0.0, 0.0], # 节点 2(噪声) [0.5, 0.0]]) # 节点 3(中等) # GCN:归一化邻接权重 A_hat = A + jnp.eye(4) D_inv = jnp.diag(1.0 / A_hat.sum(axis=1)) gcn_weights = (D_inv @ A_hat)[0] # 节点 0 的权重 print(f"GCN weights for node 0: {gcn_weights}") print(" → 所有邻居拿到大致相等的权重") # GAT:学习到的注意力(模拟) # 假设注意力机制学到要聚焦于节点 1 gat_weights = jnp.array([0.1, 0.7, 0.05, 0.15]) # 学到的 print(f"\nGAT weights for node 0: {gat_weights}") print(" → 节点 1(有信息)得到最多的注意力") gcn_output = gcn_weights @ H gat_output = gat_weights @ H print(f"\nGCN output: {gcn_output} (被噪声稀释)") print(f"GAT output: {gat_output} (聚焦于信号)")
import jax.numpy as jnp import matplotlib.pyplot as plt # 哑铃图:两个团由一座桥连接 n = 10 A = jnp.zeros((n, n)) # 团 1:节点 0-4 for i in range(5): for j in range(i+1, 5): A = A.at[i,j].set(1).at[j,i].set(1) # 团 2:节点 5-9 for i in range(5, 10): for j in range(i+1, 10): A = A.at[i,j].set(1).at[j,i].set(1) # 桥 A = A.at[4,5].set(1).at[5,4].set(1) D = jnp.diag(A.sum(axis=1)) L = D - A eigenvalues, eigenvectors = jnp.linalg.eigh(L) # 用前 3 个非平凡特征向量作为位置编码 pe = eigenvectors[:, 1:4] print("Laplacian Positional Encodings:") for i in range(n): group = "Clique 1" if i < 5 else "Clique 2" bridge = " (bridge)" if i in [4, 5] else "" print(f" Node {i} ({group}{bridge}): {pe[i]}") plt.scatter(pe[:5, 0], pe[:5, 1], c="#3498db", s=80, label="Clique 1") plt.scatter(pe[5:, 0], pe[5:, 1], c="#e74c3c", s=80, label="Clique 2") plt.scatter(pe[[4,5], 0], pe[[4,5], 1], c="black", s=120, marker="*", label="Bridge nodes", zorder=5) plt.legend(); plt.grid(True) plt.title("Laplacian Eigenvector Positional Encodings") plt.xlabel("Eigenvector 1"); plt.ylabel("Eigenvector 2") plt.show()