GAT实战:用Python从零实现图注意力网络(附完整代码)
·
GAT实战:用Python从零实现图注意力网络(附完整代码)
1. 图注意力网络的核心思想
图注意力网络(Graph Attention Network, GAT)的核心创新在于将注意力机制引入图神经网络中,使每个节点能够动态地为不同邻居分配不同的重要性权重。与传统图卷积网络(GCN)对所有邻居等权重处理的方式不同,GAT通过以下关键机制实现差异化信息聚合:
-
注意力系数计算:对于中心节点i和邻居节点j,通过可学习的权重向量a计算注意力得分:
e_ij = LeakyReLU(a^T [W h_i || W h_j])其中
||表示向量拼接,W是共享的线性变换矩阵。 -
Softmax归一化:对每个节点的邻居注意力得分进行归一化:
α_ij = softmax_j(e_ij) = exp(e_ij) / Σ_k exp(e_ik) -
多头注意力聚合:并行运行K组独立的注意力机制,增强模型表达能力:
h_i' = σ(1/K Σ_k Σ_j α_ij^k W^k h_j)
2. 环境准备与数据加载
2.1 安装依赖库
pip install torch torch-geometric ogb
2.2 加载OGBN-Arxiv数据集
from ogb.nodeproppred import PygNodePropPredDataset
dataset = PygNodePropPredDataset(name='ogbn-arxiv')
data = dataset[0] # 获取图数据
# 数据集统计信息
print(f'节点数: {data.num_nodes}')
print(f'边数: {data.num_edges}')
print(f'特征维度: {data.num_features}')
print(f'类别数: {dataset.num_classes}')
2.3 数据预处理
import torch
from torch_geometric.utils import to_undirected
# 将边索引转为无向图
data.edge_index = to_undirected(data.edge_index)
# 划分训练/验证/测试集
split_idx = dataset.get_idx_split()
train_idx = split_idx['train']
val_idx = split_idx['valid']
test_idx = split_idx['test']
3. GAT层实现详解
3.1 基础GAT层实现
import torch
import torch.nn as nn
import torch.nn.functional as F
class GATLayer(nn.Module):
def __init__(self, in_dim, out_dim, heads=1):
super().__init__()
self.heads = heads
self.out_dim = out_dim
# 线性变换矩阵
self.W = nn.Linear(in_dim, out_dim * heads, bias=False)
# 注意力参数向量
self.attn_vec = nn.Parameter(torch.randn(2 * out_dim, 1))
self.leakyrelu = nn.LeakyReLU(0.2)
def forward(self, h, adj):
# 特征投影 [N, in_dim] -> [N, heads, out_dim]
h_proj = self.W(h).view(-1, self.heads, self.out_dim)
N = h.size(0)
# 构建注意力输入 [N, heads, out_dim, N]
h_i = h_proj.repeat(1, 1, N).view(N, self.heads, self.out_dim, N)
h_j = h_proj.repeat(N, 1, 1).view(N, self.heads, self.out_dim, N)
h_cat = torch.cat([h_i, h_j], dim=2)
# 计算注意力能量
energy = torch.einsum('nhdc,cd->nhd', h_cat, self.attn_vec)
energy = self.leakyrelu(energy.squeeze(3))
# 掩码非邻居
energy = energy.masked_fill(adj == 0, -1e9)
attn = F.softmax(energy, dim=-1)
# 邻居聚合
output = torch.einsum('nhj,nhd->nhd', attn, h_proj)
return output.mean(dim=1) # 多头平均
3.2 使用PyG的高级实现
from torch_geometric.nn import GATConv
class GATLayerPyG(nn.Module):
def __init__(self, in_dim, out_dim, heads=8):
super().__init__()
self.conv = GATConv(
in_dim, out_dim, heads=heads,
dropout=0.6, concat=True
)
self.dropout = nn.Dropout(0.6)
def forward(self, x, edge_index):
x = self.dropout(x)
return self.conv(x, edge_index)
4. 完整GAT模型构建
4.1 模型架构
class GATModel(nn.Module):
def __init__(self, in_dim, hidden_dim, out_dim, heads=8):
super().__init__()
self.conv1 = GATLayerPyG(in_dim, hidden_dim, heads=heads)
self.conv2 = GATLayerPyG(hidden_dim * heads, out_dim, heads=1)
self.dropout = nn.Dropout(0.6)
def forward(self, x, edge_index):
x = self.dropout(x)
x = self.conv1(x, edge_index)
x = F.elu(x)
x = self.dropout(x)
x = self.conv2(x, edge_index)
return F.log_softmax(x, dim=1)
4.2 模型参数初始化
def init_weights(m):
if isinstance(m, nn.Linear):
nn.init.xavier_normal_(m.weight)
if m.bias is not None:
nn.init.zeros_(m.bias)
model = GATModel(
in_dim=data.num_features,
hidden_dim=256,
out_dim=dataset.num_classes
)
model.apply(init_weights)
5. 训练与评估
5.1 训练循环
optimizer = torch.optim.Adam(model.parameters(), lr=0.005, weight_decay=5e-4)
criterion = nn.NLLLoss()
def train():
model.train()
optimizer.zero_grad()
out = model(data.x, data.edge_index)
loss = criterion(out[train_idx], data.y[train_idx].view(-1))
loss.backward()
optimizer.step()
return loss.item()
@torch.no_grad()
def test():
model.eval()
out = model(data.x, data.edge_index)
# 计算各分集准确率
y_pred = out.argmax(dim=-1, keepdim=True)
correct = lambda idx: y_pred[idx].eq(data.y[idx]).sum().item() / idx.size(0)
return (
correct(train_idx),
correct(val_idx),
correct(test_idx)
)
5.2 训练过程监控
best_val_acc = 0
for epoch in range(200):
loss = train()
train_acc, val_acc, test_acc = test()
if val_acc > best_val_acc:
best_val_acc = val_acc
torch.save(model.state_dict(), 'best_model.pt')
if epoch % 10 == 0:
print(f'Epoch {epoch:03d}: Loss={loss:.4f}, '
f'Train={train_acc:.4f}, Val={val_acc:.4f}, '
f'Test={test_acc:.4f}')
6. 工程实践技巧
6.1 维度处理常见问题
-
特征维度对齐:确保每层输出的特征维度与下一层输入维度匹配,特别是多头注意力时:
# 第一层输出维度 = hidden_dim * heads # 第二层输入维度需与之匹配 -
注意力掩码实现:正确处理稀疏邻接矩阵,避免内存溢出:
# 使用稀疏矩阵运算 from torch_sparse import matmul def sparse_attention(attn, edge_index): return matmul(attn, edge_index)
6.2 性能优化策略
-
邻居采样:对于大规模图,采用邻居采样策略:
from torch_geometric.loader import NeighborSampler train_loader = NeighborSampler( data.edge_index, node_idx=train_idx, sizes=[25, 10], batch_size=1024, shuffle=True ) -
混合精度训练:使用AMP加速训练:
from torch.cuda.amp import GradScaler, autocast scaler = GradScaler() with autocast(): out = model(data.x, data.edge_index) loss = criterion(out[train_idx], data.y[train_idx]) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()
7. 进阶改进方案
7.1 残差连接
class ResidualGATLayer(nn.Module):
def __init__(self, in_dim, out_dim, heads=1):
super().__init__()
self.gat = GATLayer(in_dim, out_dim, heads)
if in_dim != out_dim:
self.residual = nn.Linear(in_dim, out_dim)
else:
self.residual = nn.Identity()
def forward(self, h, adj):
return F.elu(self.gat(h, adj) + self.residual(h))
7.2 边特征融合
class EdgeAwareGAT(nn.Module):
def __init__(self, in_dim, edge_dim, out_dim):
super().__init__()
self.edge_proj = nn.Linear(edge_dim, out_dim)
self.node_proj = nn.Linear(in_dim, out_dim)
self.attn = nn.Parameter(torch.randn(3 * out_dim, 1))
def forward(self, x, edge_index, edge_attr):
row, col = edge_index
x_i, x_j = self.node_proj(x[row]), self.node_proj(x[col])
e_ij = self.edge_proj(edge_attr)
# 融合节点和边特征
alpha = torch.cat([x_i, x_j, e_ij], dim=-1)
alpha = torch.matmul(alpha, self.attn).squeeze()
alpha = F.leaky_relu(alpha)
# 注意力加权聚合
alpha = softmax(alpha, row, num_nodes=x.size(0))
out = scatter(x_j * alpha.unsqueeze(-1), row, dim=0, dim_size=x.size(0))
return out
8. 完整代码示例
以下是一个完整的端到端GAT实现,包含数据加载、模型定义、训练和评估:
import torch
import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import GATConv
from ogb.nodeproppred import PygNodePropPredDataset
from torch_geometric.utils import to_undirected
# 数据加载
dataset = PygNodePropPredDataset(name='ogbn-arxiv')
data = dataset[0]
data.edge_index = to_undirected(data.edge_index)
split_idx = dataset.get_idx_split()
# 模型定义
class GAT(nn.Module):
def __init__(self, in_dim, hidden_dim, out_dim, heads=8):
super().__init__()
self.conv1 = GATConv(in_dim, hidden_dim, heads=heads, dropout=0.6)
self.conv2 = GATConv(hidden_dim*heads, out_dim, heads=1, concat=False, dropout=0.6)
self.dropout = nn.Dropout(0.6)
def forward(self, x, edge_index):
x = self.dropout(x)
x = self.conv1(x, edge_index)
x = F.elu(x)
x = self.dropout(x)
x = self.conv2(x, edge_index)
return F.log_softmax(x, dim=1)
# 训练设置
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = GAT(data.num_features, 256, dataset.num_classes).to(device)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
data = data.to(device)
# 训练循环
for epoch in range(200):
model.train()
optimizer.zero_grad()
out = model(data.x, data.edge_index)
loss = F.nll_loss(out[split_idx['train']], data.y[split_idx['train']].view(-1))
loss.backward()
optimizer.step()
# 验证
if epoch % 10 == 0:
model.eval()
with torch.no_grad():
pred = model(data.x, data.edge_index).argmax(dim=-1)
train_acc = pred[split_idx['train']].eq(data.y[split_idx['train']].view(-1)).float().mean()
val_acc = pred[split_idx['valid']].eq(data.y[split_idx['valid']].view(-1)).float().mean()
print(f'Epoch {epoch:03d}: Loss={loss:.4f}, Train={train_acc:.4f}, Val={val_acc:.4f}')
更多推荐
所有评论(0)