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}')
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐