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 |
dropout | — | Dropout 概率 | 0.6 |
leaky_relu_negative_slope | — | LeakyReLU 负半轴斜率 | 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=False:self.n_hidden = out_features,各头输出维度等于总输出维度,最后取平均。
随后构建五个核心模块(第 104-119 行):
self.linear_l = nn.Linear(in_features, n_hidden * n_heads, bias=False):源节点(查询侧)线性变换 $\mathbf{W}_l$;self.linear_r:目标节点(键侧)线性变换 $\mathbf{W}_r$。若share_weights=True则直接复用linear_l(self.linear_r = self.linear_l),否则新建独立矩阵;self.attn = nn.Linear(n_hidden, 1, bias=False):输出注意力分数 $e_{ij}$ 的权重向量 $\mathbf{a}$;self.activation = nn.LeakyReLU(negative_slope=leaky_relu_negative_slope):注意力打分前的非线性激活;self.softmax = nn.Softmax(dim=1):对每个查询节点的邻居做归一化,得到注意力系数 $\alpha_{ij}$;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)$ 对的注意力分数,代码用repeat与repeat_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_hidden,n_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_samples | 500 | 参与训练的节点数,其余节点用于验证 |
in_features | 由数据集计算 | 每节点输入特征数(Cora 为 1433 维词袋向量) |
n_hidden | 64 | 第一层隐藏特征数 |
n_heads | 8 | 注意力头数 |
n_classes | 由数据集计算 | 分类类别数(Cora 为 7) |
dropout | 0.6 | Dropout 概率 |
include_edges | True | 是否使用引用边(设为False可测试丢掉图结构后的精度损失) |
epochs | 1_000 | 训练迭代轮数 |
loss_func | nn.CrossEntropyLoss() | 分类损失函数 |
device | DeviceConfigs() | 训练设备(可通过配置切换) |
optimizer | — | 通过OptimizerConfigs可配置化的优化器 |
n_classes与in_features由calculate装饰器根据数据集自动推导(见 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] = True与adj_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.train、loss.valid、accuracy.train、accuracy.valid指标。
七、运行训练
仓库使用 labml 实验框架,安装依赖后即可直接运行。依赖清单见 requirements.txt(核心为torch>=1.10与labml>=0.4.147),也可通过setup.py以pip 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 算子实现要点对比
| 对比维度 | 标准 GAT | GATv2 |
|---|---|---|
| 注意力公式 | $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) |
| 激活位置 | 打分之后再激活 | 打分之前激活(关键差异) |
| 实验默认 Dropout | 0.6 | 0.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),仅供参考