GATv2 图注意力网络实现与源码解析:从静态注意力缺陷到 Cora 节点分类实战
2026/9/18 5:11:04 网站建设 项目流程

GATv2 图注意力网络实现与源码解析:从静态注意力缺陷到 Cora 节点分类实战

【免费下载链接】annotated_deep_learning_paper_implementations🧑‍🏫 60+ Implementations/tutorials of deep learning papers with side-by-side notes 📝; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), 🎮 reinforcement learning (ppo, dqn), capsnet, distillation, ... 🧠项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations

本篇技术指南围绕 annotated_deep_learning_paper_implementations 仓库中的 GATv2 文档 展开,系统讲解 Graph Attention Networks v2(GATv2)算子的设计动机、数学原理,并结合仓库内的 PyTorch 核心实现 与 Cora 训练代码 逐行剖析其前向传播细节与实验配置。读完本文,你将掌握 GATv2 相对标准 GAT 的关键改进点(动态注意力机制),能够独立理解、复用并扩展该算子,并能在 Cora 引文网络上复现两层 GATv2 的节点分类训练。

一、GATv2 是什么:面向图数据的注意力算子

GATv2 是论文How Attentive are Graph Attention Networks?(arXiv:2105.14491)提出的图注意力算子,本仓库以 PyTorch 完整实现,相关代码与注释同时承担"实现 + 教学"双重角色。

GATv2 面向图数据(graph data)工作。所谓图,由一组**节点(nodes)和连接节点的边(edges)**构成。以 Cora 数据集为例:图中的节点是研究论文,边则是论文之间的引用关系——论文 A 引用论文 B,就在两篇论文对应的节点间建立一条边。图神经网络的任务,就是利用节点自身的特征向量以及节点之间的连接结构,学习每个节点的高质量表示。

在仓库结构中,图神经网络相关实现统一放在 labml_nn/graphs/ 目录下,其中gat/存放标准 GAT,gatv2/存放 GATv2 的算子实现、训练代码与本文所依据的 readme 文档:

  • labml_nn/graphs/gatv2/init.py:GraphAttentionV2Layer单层算子实现;
  • labml_nn/graphs/gatv2/experiment.py:在 Cora 数据集上训练两层 GATv2 的完整实验代码;
  • labml_nn/graphs/gat/init.py 与 labml_nn/graphs/gat/experiment.py:标准 GAT 的算子与训练代码,用作对比参照。

二、GATv2 要解决的问题:标准 GAT 的"静态注意力"缺陷

GATv2 的核心贡献是修复标准 GAT 的静态注意力(static attention)问题。为了理解这一点,先看标准 GAT 的注意力分数计算方式。

2.1 标准 GAT 的注意力公式

标准 GAT 计算从查询节点 $i$ 到键节点 $j$ 的注意力分数 $e_{ij}$ 时,先对源节点和目标节点使用同一个线性变换 $\mathbf{W}$,再将两者拼接后与注意力向量 $\mathbf{a}$ 做内积,最后经过 LeakyReLU 激活:

$$ \begin{align} e_{ij} &= \text{LeakyReLU} \Big(\mathbf{a}^\top \Big[ \mathbf{W} \overrightarrow{h_i} \Vert \mathbf{W} \overrightarrow{h_j} \Big] \Big) \ &= \text{LeakyReLU} \Big(\mathbf{a}_1^\top \mathbf{W} \overrightarrow{h_i} + \mathbf{a}_2^\top \mathbf{W} \overrightarrow{h_j} \Big) \end{align} $$

展开后可以看出,拼接等价于把注意力向量 $\mathbf{a}$ 拆成 $\mathbf{a}_1$ 与 $\mathbf{a}_2$ 两部分,分别作用于 $\mathbf{W}\overrightarrow{h_i}$ 与 $\mathbf{W}\overrightarrow{h_j}$,再求和。

2.2 缺陷所在:注意力排名与查询节点无关

关键观察在于:对任意查询节点 $i$,各键节点的注意力排名(即对 $e_{ij}$ 做 $argsort$ 得到的次序)只取决于 $\mathbf{a}_2^\top \mathbf{W} \overrightarrow{h_j}$ 这一项

因为 $\mathbf{a}1^\top \mathbf{W} \overrightarrow{h_i}$ 只与查询节点 $i$ 有关,它作为常数项加到所有 $e{ij}$ 上,不会改变排序。也就是说,无论查询节点是谁,被关注节点(键)的相对顺序始终保持一致。这种注意力被称为静态注意力——模型对所有查询都"一视同仁"地按同一顺序关注邻居,无法针对不同的查询节点差异化地重新排列关注对象。

2.3 静态注意力的失败场景

论文用一个合成的字典查找(dictionary lookup)数据集展示了 GAT 的失败:这是一个全连接的二分图,一侧是查询节点(query nodes),每个查询节点关联一个 key;另一侧节点同时关联一个 key 和一个 value。任务是根据查询节点的 key,预测其对应 value。由于注意力是静态的,GAT 无法针对不同查询灵活调整关注目标,在这类任务上表现不佳。

三、GATv2 的动态注意力:改变算子运算次序

GATv2 的改进思路非常直观:交换线性变换与激活函数的次序,让查询节点参与非线性变换之后再打分。其注意力分数定义为:

$$ \begin{align} e_{ij} &= \mathbf{a}^\top \text{LeakyReLU} \Big( \mathbf{W} \Big[ \overrightarrow{h_i} \Vert \overrightarrow{h_j} \Big] \Big) \ &= \mathbf{a}^\top \text{LeakyReLU} \Big( \mathbf{W}_l \overrightarrow{h_i} + \mathbf{W}_r \overrightarrow{h_j} \Big) \end{align} $$

对比标准 GAT:GAT 是"线性变换 → 拼接 → 内积 → 激活",且源与目标共享同一个 $\mathbf{W}$;GATv2 是"分别线性变换 → 相加 → 激活 → 内积",并使用两个不同的矩阵 $\mathbf{W}_l$ 与 $\mathbf{W}_r$。

由于 LeakyReLU 是非线性的,$\mathbf{a}^\top \text{LeakyReLU}(\mathbf{W}_l \overrightarrow{h_i} + \mathbf{W}_r \overrightarrow{h_j})$ 无法再分解为"与 $i$ 无关的项 + 与 $j$ 无关的项",因此注意力排名真正同时依赖查询节点与键节点——每个节点都可以关注到任意其他节点,注意力变为动态的(dynamic attention)。

更直白的等价理解:GATv2 相当于先对 $\mathbf{W}_l \overrightarrow{h_i} + \mathbf{W}_r \overrightarrow{h_j}$ 施加一个非线性变换,再用向量 $\mathbf{a}$ 度量其"方向",其表达能力显著强于 GAT 中"拼接后线性打分"的线性度量。

四、GraphAttentionV2Layer 源码级解析

仓库的核心算子实现位于 labml_nn/graphs/gatv2/init.py,核心类为GraphAttentionV2Layer(定义于 第 62 行)。一个 GATv2 网络由多个这样的层堆叠而成:每层接收节点嵌入集合 $\mathbf{h} = { \overrightarrow{h_1}, \overrightarrow{h_2}, \dots, \overrightarrow{h_N} }$(其中 $\overrightarrow{h_i} \in \mathbb{R}^F$)作为输入,输出变换后的嵌入集合 $\mathbf{h'} = { \overrightarrow{h'_1}, \dots, \overrightarrow{h'_N} }$(其中 $\overrightarrow{h'_i} \in \mathbb{R}^{F'}$)。

4.1 构造参数与内部模块

__init__(第 75-119 行)接收以下参数:

参数符号含义默认值
in_features$F$每个节点的输入特征数必填
out_features$F'$每个节点的输出特征数必填
n_heads$K$注意力头数必填
is_concat多头结果是拼接还是取平均True
dropoutDropout 概率0.6
leaky_relu_negative_slopeLeakyReLU 负半轴斜率0.2
share_weights源节点与目标节点是否共享线性层False

初始化逻辑中,先根据是否拼接多头来确定每个头内部的隐藏维度:

  • is_concat=True:断言out_features % n_heads == 0,令self.n_hidden = out_features // n_heads,即 $F' = K \times F'_{head}$;
  • is_concat=Falseself.n_hidden = out_features,各头输出维度等于总输出维度,最后取平均。

随后构建五个核心模块(第 104-119 行):

  1. self.linear_l = nn.Linear(in_features, n_hidden * n_heads, bias=False):源节点(查询侧)线性变换 $\mathbf{W}_l$;
  2. self.linear_r:目标节点(键侧)线性变换 $\mathbf{W}_r$。若share_weights=True则直接复用linear_lself.linear_r = self.linear_l),否则新建独立矩阵;
  3. self.attn = nn.Linear(n_hidden, 1, bias=False):输出注意力分数 $e_{ij}$ 的权重向量 $\mathbf{a}$;
  4. self.activation = nn.LeakyReLU(negative_slope=leaky_relu_negative_slope):注意力打分前的非线性激活;
  5. self.softmax = nn.Softmax(dim=1):对每个查询节点的邻居做归一化,得到注意力系数 $\alpha_{ij}$;
  6. self.dropout = nn.Dropout(dropout):对注意力系数施加正则。

注意一个细节差异:标准 GAT 中attn层的输入维度是n_hidden * 2(拼接两个节点变换后的向量),而 GATv2 中是n_hidden(先相加再打分),这正是两个算子结构差异在代码层面的直接体现。

4.2 前向传播的五步流水线

forward(第 121 行起)接收两个张量:

  • h:节点嵌入,形状[n_nodes, in_features]
  • adj_mat:邻接矩阵,形状[n_nodes, n_nodes, n_heads][n_nodes, n_nodes, 1](仓库实现中因各头共享同一邻接结构而使用单通道);adj_mat[i][j]True表示存在从节点i到节点j的边。

第一步:初始双线性变换并切分多头(第 131-138 行)

对每个头 $k$ 计算 $\overrightarrow{{g_l}^k_i} = \mathbf{W_l}^k \overrightarrow{h_i}$ 与 $\overrightarrow{{g_r}^k_i} = \mathbf{W_r}^k \overrightarrow{h_i}$:

g_l = self.linear_l(h).view(n_nodes, self.n_heads, self.n_hidden) g_r = self.linear_r(h).view(n_nodes, self.n_heads, self.n_hidden)

第二步:构造所有节点对的组合(第 168-191 行)

为了一次性计算所有 $(i, j)$ 对的注意力分数,代码用repeatrepeat_interleave构造笛卡尔积:

  • g_l_repeat = g_l.repeat(n_nodes, 1, 1):把每个节点的嵌入整体重复n_nodes次,得到序列 ${\overrightarrow{{g_l}_1}, \dots, \overrightarrow{{g_l}_N}, \overrightarrow{{g_l}_1}, \dots, \overrightarrow{{g_l}_N}, \dots}$;
  • g_r_repeat_interleave = g_r.repeat_interleave(n_nodes, dim=0):把每个节点的嵌入逐份交错重复n_nodes次,得到 ${\overrightarrow{{g_r}_1}, \overrightarrow{{g_r}_1}, \dots, \overrightarrow{{g_r}_N}, \dots}$;

两者相加再view(n_nodes, n_nodes, n_heads, n_hidden)后,g_sum[i, j]正好等于 $\overrightarrow{{g_l}_i} + \overrightarrow{{g_r}_j}$。

第三步:计算注意力分数 $e_{ij}$(第 193-201 行)

$$ e_{ij} = \mathbf{a}^\top \text{LeakyReLU} \Big( \overrightarrow{{g_l}_i} + \overrightarrow{{g_r}_j} \Big) $$

e = self.attn(self.activation(g_sum)) # [n_nodes, n_nodes, n_heads, 1] e = e.squeeze(-1) # [n_nodes, n_nodes, n_heads]

这里正是 GATv2 与 GAT 的分水岭:GAT 是"先拼接、再线性、后激活"(e = self.activation(self.attn(g_concat)),见 GAT 实现);GATv2 是"先相加、再激活、后线性"。

第四步:邻接掩码与 softmax 归一化(第 203-223 行)

先通过断言校验邻接矩阵形状合法性,然后:

e = e.masked_fill(adj_mat == 0, float('-inf')) a = self.softmax(e) a = self.dropout(a)

将不存在的边对应的 $e_{ij}$ 置为 $-\infty$,使 $\exp(e_{ij}) \approx 0$,从而 softmax 只在邻居集合 $\mathcal{N}_i$ 上归一化:

$$ \alpha_{ij} = \text{softmax}j(e{ij}) = \frac{\exp(e_{ij})}{\sum_{j' \in \mathcal{N}i} \exp(e{ij'})} $$

随后对注意力系数施加 Dropout 正则。

第五步:加权聚合与多头合并(第 225-236 行)

用 einsum 完成对目标节点变换结果的加权求和:

attn_res = torch.einsum('ijh,jhf->ihf', a, g_r)

即 $\overrightarrow{h'^k_i} = \sum_{j \in \mathcal{N}i} \alpha^k{ij} \overrightarrow{{g_r}_{j,k}}$。最后按is_concat决定输出方式:

  • 拼接多头:$\overrightarrow{h'i} = \Bigg\Vert{k=1}^{K} \overrightarrow{h'^k_i}$,形状变为n_nodes * n_heads * n_hidden
  • 取平均:$\overrightarrow{h'i} = \frac{1}{K} \sum{k=1}^{K} \overrightarrow{h'^k_i}$。

五、两层 GATv2 网络与训练配置

5.1 模型结构

experiment.py 中的GATv2类(第 20-65 行)构造了一个两层 GATv2:

self.layer1 = GraphAttentionV2Layer(in_features, n_hidden, n_heads, is_concat=True, dropout=dropout, share_weights=share_weights) self.activation = nn.ELU() self.output = GraphAttentionV2Layer(n_hidden, n_classes, 1, is_concat=False, dropout=dropout, share_weights=share_weights) self.dropout = nn.Dropout(dropout)
  • 第一层in_features → n_hiddenn_heads个头且结果拼接is_concat=True);
  • 中间激活用nn.ELU(),并在输入层和激活之后各施加一次 Dropout;
  • 输出层n_hidden → n_classes,仅 1 个头且结果取平均is_concat=False),不接激活函数,直接输出分类 logits。

forward的输入约定与单层算子一致:x形状为[n_nodes, in_features]adj_mat形状为[n_nodes, n_nodes, n_heads][n_nodes, n_nodes, 1]

5.2 配置继承与覆盖

Configs类(第 68-79 行)直接继承标准 GAT 实验的配置类GATConfigs(定义于 GAT 训练代码),因为两者的实验框架几乎完全一致,只需替换模型:

class Configs(GATConfigs): # 源节点与目标节点是否共享权重矩阵 share_weights: bool = False # 将模型切换为 GATv2 model: GATv2 = 'gat_v2_model'

share_weights默认为False(源、目标各用各的 $\mathbf{W}_l$、$\mathbf{W}_r$),这是论文推荐设置的默认行为;若设为True,则复用同一个矩阵(对应论文中的共享权重变体)。

从父类继承的关键配置项(均可在运行前覆盖)包括:

配置项默认值说明
training_samples500参与训练的节点数,其余节点用于验证
in_features由数据集计算每节点输入特征数(Cora 为 1433 维词袋向量)
n_hidden64第一层隐藏特征数
n_heads8注意力头数
n_classes由数据集计算分类类别数(Cora 为 7)
dropout0.6Dropout 概率
include_edgesTrue是否使用引用边(设为False可测试丢掉图结构后的精度损失)
epochs1_000训练迭代轮数
loss_funcnn.CrossEntropyLoss()分类损失函数
deviceDeviceConfigs()训练设备(可通过配置切换)
optimizer通过OptimizerConfigs可配置化的优化器

n_classesin_featurescalculate装饰器根据数据集自动推导(见 GAT 训练代码),无需手动指定。

5.3 实验入口与超参数

main()(第 90-108 行)创建名为gatv2的实验并设置优化器与正则化超参数:

experiment.create(name='gatv2') experiment.configs(conf, { 'optimizer.optimizer': 'Adam', 'optimizer.learning_rate': 5e-3, 'optimizer.weight_decay': 5e-4, 'dropout': 0.7, })

即使用Adam 优化器,学习率 5e-3,权重衰减 5e-4,并将 Dropout 覆盖为0.7(比父类默认的 0.6 更强)。模型通过gat_v2_model工厂函数创建并移动到指定设备(第 82-87 行)。

六、Cora 数据集与训练循环

6.1 Cora 数据加载

Cora 数据集的加载逻辑定义在 GAT 训练代码 的CoraDataset类中,GATv2 实验直接复用:

  • 数据内容cora.content文件为每篇论文提供二值词袋特征向量与 7 个类别之一的标签;cora.cites文件记录论文间的引用对;
  • 特征预处理:特征向量按行归一化(features / features.sum(dim=1, keepdim=True)),features.shape[1]in_features
  • 类别映射:类别名映射为唯一整数索引,类别数即n_classes
  • 邻接矩阵:初始化为单位阵(每个节点自带自环),随后对每条引用边(e1, e2)同时置adj_mat[e1][e2] = Trueadj_mat[e2][e1] = True,构建对称无向图
  • 自动下载:首次运行自动下载并解压数据集到 labml 数据目录(由_download方法完成)。

6.2 训练循环

Configs.run()(第 194-253 行)实现完整的训练流程,关键点:

  • 全批量训练:Cora 数据集较小,因此不采样、直接对全图训练;注释说明若改为采样训练,还需同时采样跨越所选节点的边;
  • 数据划分:用torch.randperm随机打乱节点索引,前training_samples(500)个节点作为训练集,其余作为验证集;
  • 单步流程optimizer.zero_grad()→ 前向得到全图 logits → 只对训练节点计算交叉熵损失 →loss.backward()optimizer.step()
  • 验证评估:切换model.eval(),在torch.no_grad()下对验证节点计算损失与准确率,并通过tracker记录loss.trainloss.validaccuracy.trainaccuracy.valid指标。

七、运行训练

仓库使用 labml 实验框架,安装依赖后即可直接运行。依赖清单见 requirements.txt(核心为torch>=1.10labml>=0.4.147),也可通过setup.pypip install -e .方式安装本项目。

训练两层 GATv2 于 Cora 数据集:

python -m labml_nn.graphs.gatv2.experiment

运行后会自动下载 Cora 数据集,创建名为gatv2的实验,按上述配置训练 1000 轮并记录损失与准确率指标。若想对比标准 GAT 在同一数据集上的表现,可运行python -m labml_nn.graphs.gat.experiment(对应 GAT 训练代码)。

八、GAT 与 GATv2 算子实现要点对比

对比维度标准 GATGATv2
注意力公式$e_{ij} = \text{LeakyReLU}(\mathbf{a}^\top[\mathbf{W}h_i \Vert \mathbf{W}h_j])$$e_{ij} = \mathbf{a}^\top \text{LeakyReLU}(\mathbf{W}_l h_i + \mathbf{W}_r h_j)$
源/目标变换共享同一 $\mathbf{W}$默认两个矩阵 $\mathbf{W}_l$、$\mathbf{W}_r$(可选共享)
注意力类型静态:键节点排名与查询无关动态:排名同时依赖查询与键
拼接 vs 相加变换后拼接(attn输入维2 * n_hidden变换后相加(attn输入维n_hidden
激活位置打分之后再激活打分之前激活(关键差异)
实验默认 Dropout0.60.7

从源码结构可以推断:GATv2 的改动虽然只涉及算子内部的运算次序与权重拆分,却从根本上改变了注意力的表达能力——这正是"动态注意力"得以成立的关键。两个实验共用同一套 Cora 数据加载、训练循环与配置体系,唯一的差异在于算子与模型,因此 GATv2 实验 通过继承 GAT 实验 并覆盖model配置的方式,最小化地复用了全部基础设施,是仓库中"以配置切换模型"的典型范例。

总结

GATv2 以一处精巧的公式重排解决了标准 GAT 的静态注意力缺陷:将 LeakyReLU 从打分之后移到打分之前,配合源/目标分离的线性变换,使得注意力分数能够真正随查询节点变化。仓库中的 单层算子实现 以可读的教学式注释完整呈现了"双线性变换 → 笛卡尔积组合 → 激活打分 → 邻接掩码 → softmax → einsum 聚合 → 多头合并"的全流程,而 实验代码 则展示了如何复用 GAT 的配置与数据管线,在 Cora 引文网络上直接训练一个两层 GATv2 完成论文分类任务。无论你是要复现论文、对比 GAT 变体,还是将其作为组件接入自己的图学习项目,这份实现都可以作为直接可用的起点。

【免费下载链接】annotated_deep_learning_paper_implementations🧑‍🏫 60+ Implementations/tutorials of deep learning papers with side-by-side notes 📝; including transformers (original, xl, switch, feedback, vit, ...), optimizers (adam, adabelief, sophia, ...), gans(cyclegan, stylegan2, ...), 🎮 reinforcement learning (ppo, dqn), capsnet, distillation, ... 🧠项目地址: https://gitcode.com/gh_mirrors/an/annotated_deep_learning_paper_implementations

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询