GraphSAGE实战:用PyTorch Geometric快速搭建图神经网络(附完整代码)

最近几年,图神经网络(GNN)的热度居高不下,从学术论文到工业级应用,处处都能看到它的身影。但很多刚接触这个领域的朋友,包括我自己刚开始的时候,都会觉得有点无从下手——理论听起来很美好,但代码怎么跑起来?模型怎么调?数据怎么处理?如果你也有类似的困惑,那这篇文章就是为你准备的。我们不打算从复杂的数学公式讲起,而是直接切入实战,用PyTorch Geometric这个强大的库,手把手带你搭建一个可运行的GraphSAGE模型。无论你是想快速验证一个想法,还是需要在项目中集成图学习能力,这篇指南都能帮你节省大量摸索的时间。

1. 环境准备与数据理解

在开始写代码之前,我们需要把“战场”打扫干净。对于图神经网络项目,环境配置和数据理解是两大基石,这一步走稳了,后面的开发会顺畅很多。

1.1 搭建你的开发环境

我强烈建议使用Anaconda来管理Python环境,它能很好地解决不同项目间的依赖冲突问题。首先,创建一个新的虚拟环境:

conda create -n graphsage_env python=3.9
conda activate graphsage_env

接下来安装核心的PyTorch。请务必根据你的CUDA版本(如果有GPU的话)去PyTorch官网选择对应的安装命令。例如,对于CUDA 11.8,可以这样安装:

pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118

安装完PyTorch,就可以安装我们今天的主角——PyTorch Geometric(简称PyG)。由于它依赖一些特殊的扩展库(如torch-scatter, torch-sparse),直接pip install可能会失败。最稳妥的方式是使用PyG官方提供的预编译包。假设你安装的是PyTorch 2.0+和CUDA 11.8,命令如下:

pip install pyg-lib torch-scatter torch-sparse torch-cluster torch-spline-conv -f https://data.pyg.org/whl/torch-2.0.0+cu118.html
pip install torch-geometric

注意:上述URL中的torch-2.0.0+cu118需要替换为你实际安装的PyTorch版本和CUDA版本。如果使用CPU,则对应cpu版本。

安装完成后,可以写一个简单的脚本来验证环境:

import torch
import torch_geometric
print(f"PyTorch version: {torch.__version__}")
print(f"PyG version: {torch_geometric.__version__}")
print(f"CUDA available: {torch.cuda.is_available()}")

1.2 图数据:从抽象到PyG的Data对象

图数据和我们熟悉的表格数据、图像数据很不一样。它由两部分核心构成:节点。每个节点可能有自己的特征(比如社交网络中用户的年龄、性别),每条边则代表了节点间的关系。

在PyG中,图数据被封装在torch_geometric.data.Data对象里。理解这个对象的几个关键属性至关重要:

  • x (节点特征矩阵): 形状为 [num_nodes, num_node_features] 的Tensor。num_nodes是节点总数,num_node_features是每个节点的特征维度。
  • edge_index (边索引): 形状为 [2, num_edges] 的LongTensor。它定义了图中所有的连接关系。第一行是源节点索引,第二行是目标节点索引。例如,edge_index = torch.tensor([[0, 1, 1, 2], [1, 0, 2, 1]]) 表示有四条边:0->1, 1->0, 1->2, 2->1。
  • y (节点或图的标签): 根据任务不同,可以是节点级标签(形状[num_nodes])或图级标签(形状[1][num_graphs])。
  • edge_attr (边特征): 可选,形状为 [num_edges, num_edge_features]

为了让你有更直观的感受,我们用一个经典的学术数据集——Cora(引文网络)来举例。你可以通过PyG内置的数据集轻松加载它:

from torch_geometric.datasets import Planetoid

dataset = Planetoid(root='/tmp/Cora', name='Cora')
data = dataset[0]  # Cora数据集只有一个图

print(f'Number of nodes: {data.num_nodes}')
print(f'Number of edges: {data.num_edges}')
print(f'Number of node features: {data.num_node_features}')
print(f'Number of classes: {data.num_classes}')
print(f'Has isolated nodes: {data.has_isolated_nodes()}')
print(f'Has self-loops: {data.has_self_loops()}')
print(f'Is undirected: {data.is_undirected()}')

运行这段代码,你会看到Cora数据集包含2708篇论文(节点),每篇论文有1433维的特征(词袋模型),它们之间有10556条引用关系(边),任务是将论文分为7个类别。data对象还自动包含了train_mask, val_mask, test_mask,用于划分训练、验证和测试集。

2. 构建你的第一个GraphSAGE模型

理解了数据,我们就可以开始搭建模型了。GraphSAGE的核心思想是通过采样邻居并聚合其信息,来迭代地更新节点表示。PyG已经为我们实现了SAGEConv层,这大大简化了我们的工作。

2.1 模型架构设计

一个典型的GraphSAGE模型由多层SAGEConv堆叠而成。每一层都执行以下操作:

  1. 对每个节点,从其邻居中采样固定数量的节点(在SAGEConv内部默认使用所有邻居,但我们可以自定义采样器)。
  2. 将采样到的邻居节点的特征聚合起来(比如取均值、最大值或通过一个神经网络)。
  3. 将聚合后的邻居特征与节点自身特征结合,经过一个非线性变换(如ReLU),得到该节点在这一层的新特征。

下面是一个基础的GraphSAGE模型类定义:

import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import SAGEConv

class GraphSAGE(nn.Module):
    def __init__(self, in_channels, hidden_channels, out_channels, num_layers, dropout=0.5):
        super(GraphSAGE, self).__init__()
        self.convs = nn.ModuleList()
        self.dropout = dropout
        
        # 第一层:从输入维度到隐藏维度
        self.convs.append(SAGEConv(in_channels, hidden_channels))
        # 中间层:保持隐藏维度不变
        for _ in range(num_layers - 2):
            self.convs.append(SAGEConv(hidden_channels, hidden_channels))
        # 最后一层:从隐藏维度到输出维度(如类别数)
        self.convs.append(SAGEConv(hidden_channels, out_channels))

    def forward(self, x, edge_index):
        for i, conv in enumerate(self.convs[:-1]):
            x = conv(x, edge_index)  # 图卷积操作
            x = F.relu(x)            # 非线性激活
            x = F.dropout(x, p=self.dropout, training=self.training) # Dropout正则化
        # 最后一层通常不加激活函数,直接输出logits
        x = self.convs[-1](x, edge_index)
        return x

这个模型结构非常清晰。num_layers控制了网络的深度,也决定了每个节点能聚合多少“跳”邻居的信息。层数太少,模型可能无法捕获足够的全局结构;层数太多,则可能导致过平滑问题,即所有节点的表示变得过于相似。

2.2 聚合函数的选择与影响

SAGEConv默认使用均值聚合器,即对邻居特征取平均。这是最常用也最稳定的选择。但PyG的SAGEConv也支持其他聚合方式,例如max聚合。你可以在初始化时通过aggr参数指定:

# 使用最大池化聚合邻居信息
conv_layer = SAGEConv(in_channels, out_channels, aggr='max')

不同的聚合器有不同的特性:

  • 均值聚合器 (mean): 平滑、稳定,能反映邻居群体的整体趋势。
  • 最大池化聚合器 (max): 更具“判别性”,只关注邻居中最显著的特征,可能对异常值更敏感。
  • LSTM聚合器: 理论上能考虑邻居的顺序,但计算成本高且需要对邻居进行排序,实际中使用较少。

对于大多数入门和中级任务,使用默认的mean聚合器即可。当你发现模型对某些局部结构不敏感时,可以尝试max聚合器。

3. 模型训练、验证与调优实战

模型搭好了,接下来就是让它“学习”。图神经网络的训练循环和普通神经网络大同小异,但有一些细节需要特别注意。

3.1 编写训练与评估循环

我们将使用Cora数据集进行节点分类任务。完整的训练流程代码如下:

import torch.optim as optim
from torch_geometric.datasets import Planetoid
from torch_geometric.transforms import NormalizeFeatures

# 1. 加载并预处理数据
dataset = Planetoid(root='/tmp/Cora', name='Cora', transform=NormalizeFeatures())
data = dataset[0]
data = data.to(device) # 如果有GPU,将数据移到GPU上

# 2. 初始化模型、优化器和损失函数
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
model = GraphSAGE(in_channels=dataset.num_node_features,
                  hidden_channels=64,
                  out_channels=dataset.num_classes,
                  num_layers=3,
                  dropout=0.5).to(device)
optimizer = optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
criterion = nn.CrossEntropyLoss()

# 3. 训练函数
def train():
    model.train()
    optimizer.zero_grad()
    out = model(data.x, data.edge_index)          # 前向传播
    loss = criterion(out[data.train_mask], data.y[data.train_mask]) # 仅计算训练节点的损失
    loss.backward()                               # 反向传播
    optimizer.step()                              # 更新参数
    return loss.item()

# 4. 测试函数
@torch.no_grad()
def test():
    model.eval()
    out = model(data.x, data.edge_index)
    pred = out.argmax(dim=1)  # 取概率最大的类别作为预测
    
    # 分别计算训练集、验证集、测试集的准确率
    accs = []
    for mask in [data.train_mask, data.val_mask, data.test_mask]:
        correct = (pred[mask] == data.y[mask]).sum().item()
        acc = correct / mask.sum().item()
        accs.append(acc)
    return accs

# 5. 训练循环
for epoch in range(1, 201):
    loss = train()
    train_acc, val_acc, test_acc = test()
    if epoch % 20 == 0:
        print(f'Epoch: {epoch:03d}, Loss: {loss:.4f}, '
              f'Train Acc: {train_acc:.4f}, Val Acc: {val_acc:.4f}, Test Acc: {test_acc:.4f}')

这段代码有几个关键点:

  • 数据划分:Cora数据已经内置了掩码(train_mask等),我们只在训练节点上计算损失和反向传播。
  • 评估模式:在测试时使用@torch.no_grad()model.eval()来关闭Dropout和自动求导,节省内存并保证结果一致性。
  • 早停策略:代码中没有体现,但在实际中,你应该监控验证集准确率,当其在连续多个epoch不再提升时停止训练,以防止过拟合。

3.2 超参数调优策略

GraphSAGE的性能对超参数比较敏感。下面是一个简单的调优指南,你可以从这些默认值开始尝试:

超参数常见范围/选择影响与调优建议
隐藏层维度64, 128, 256维度越大,模型容量越高,但也更容易过拟合。从64或128开始。
网络层数2, 3, 4层数决定感受野。Cora这类小图2-3层足够;大规模图可能需要更多层,但要警惕过平滑。
学习率0.01, 0.005, 0.001图神经网络通常使用较小的学习率。0.01是个不错的起点。
Dropout率0.5, 0.6防止过拟合的有效正则化手段。在0.5附近调整。
权重衰减5e-4L2正则化系数,帮助控制模型复杂度。

一个实用的调优流程是:

  1. 固定其他,调整层数:先尝试2层和3层,看哪个在验证集上更好。
  2. 调整隐藏层维度:在选定层数后,尝试64和128。
  3. 微调正则化:调整Dropout和权重衰减,在训练集和验证集准确率之间寻找平衡。
  4. 降低学习率:如果训练后期损失震荡或无法收敛,可以尝试将学习率减半。

你可以使用torch.optim.lr_scheduler中的学习率调度器,比如ReduceLROnPlateau,让训练过程更自动化。

4. 进阶技巧与实战陷阱规避

当你跑通第一个模型后,可能会遇到准确率不高、训练不稳定或想应用到更复杂场景的问题。这部分分享一些我实践中总结的进阶技巧。

4.1 处理过拟合与过平滑

这是图神经网络的两个常见“杀手”。

  • 过拟合:模型在训练集上表现很好,在验证/测试集上很差。

    • 对策:加大Dropout率;增强权重衰减;使用更早的早停;如果数据允许,可以尝试数据增强(如对边进行随机丢弃,即DropEdge)。
    • 代码示例 - 边丢弃
      from torch_geometric.transforms import DropEdge
      transform = DropEdge(p=0.2) # 以20%的概率随机丢弃边
      data_transformed = transform(data)
      # 注意:需要在每个训练epoch前对数据进行变换
      
  • 过平滑:随着网络层数增加,所有节点的特征表示变得过于相似,导致分类性能下降。

    • 对策:这是GNN的固有问题。最有效的方法是不要堆叠太多层。对于类似Cora的小图,2-3层通常是上限。也可以尝试使用残差连接跳跃连接,将浅层特征传递到深层。
    • 代码示例 - 添加残差连接
      class GraphSAGEWithResidual(nn.Module):
          def __init__(self, in_channels, hidden_channels, out_channels, num_layers):
              super().__init__()
              self.convs = nn.ModuleList()
              self.convs.append(SAGEConv(in_channels, hidden_channels))
              for _ in range(num_layers - 2):
                  self.convs.append(SAGEConv(hidden_channels, hidden_channels))
              self.convs.append(SAGEConv(hidden_channels, out_channels))
              self.res_fc = nn.Linear(in_channels, hidden_channels) # 用于调整输入维度
      
          def forward(self, x, edge_index):
              h = x
              for i, conv in enumerate(self.convs[:-1]):
                  h_new = conv(h, edge_index)
                  if i == 0: # 第一层,需要将原始输入x变换到hidden维度
                      h = F.relu(h_new + self.res_fc(x))
                  else:
                      h = F.relu(h_new + h) # 残差连接
                  h = F.dropout(h, p=0.5, training=self.training)
              h = self.convs[-1](h, edge_index)
              return h
      

4.2 扩展到大规模图与自定义数据集

Cora数据集很小,可以全图加载到内存。但工业级的图往往有数百万甚至数十亿个节点。这时就需要使用邻居采样

PyG提供了NeighborLoader等工具,它会在每个训练批次中,为一批“种子节点”递归地采样其多跳邻居,构建一个用于计算的小子图,从而极大节省内存。

from torch_geometric.loader import NeighborLoader

# 假设 `data` 是一个大图的Data对象
train_loader = NeighborLoader(
    data,
    num_neighbors=[10, 5],  # 第一层采样10个邻居,第二层采样5个邻居
    batch_size=32,
    input_nodes=data.train_mask,  # 只对训练节点进行采样
    shuffle=True
)

# 训练循环需要相应调整
for batch in train_loader:
    batch = batch.to(device)
    optimizer.zero_grad()
    out = model(batch.x, batch.edge_index)
    loss = criterion(out[batch.train_mask], batch.y[batch.train_mask])
    loss.backward()
    optimizer.step()

对于自定义数据集,你需要将你的图数据构造成PyG的Data对象。常见的来源是CSV文件或NetworkX图:

import pandas as pd
import networkx as nx
from torch_geometric.utils import from_networkx

# 从CSV构建(假设有边列表文件edge_list.csv)
df_edges = pd.read_csv('edge_list.csv')
edge_index = torch.tensor([df_edges['src'].values, df_edges['dst'].values], dtype=torch.long)

# 从节点特征文件构建
df_features = pd.read_csv('node_features.csv')
x = torch.tensor(df_features.values, dtype=torch.float)

# 创建Data对象
data = Data(x=x, edge_index=edge_index)

# 或者从NetworkX图转换
G = nx.read_edgelist('graph.edgelist')
data = from_networkx(G)
# 然后需要手动添加data.x节点特征

4.3 模型保存、加载与推理

训练好的模型需要保存下来供后续使用。

# 保存整个模型
torch.save(model.state_dict(), 'graphsage_cora.pth')

# 加载模型进行推理
model = GraphSAGE(...).to(device)
model.load_state_dict(torch.load('graphsage_cora.pth'))
model.eval()

with torch.no_grad():
    out = model(data.x, data.edge_index)
    predictions = out.argmax(dim=1)
    # 对predictions进行后续处理...

最后,我想说的是,图神经网络的实践是一个不断迭代和试错的过程。第一次运行可能效果不理想,这非常正常。多动手修改代码,调整参数,观察训练曲线,并尝试理解模型在每一层学到了什么(可以通过可视化中间层特征)。GraphSAGE只是一个起点,PyG中还有GAT、GIN等更多强大的图卷积算子等着你去探索。当你掌握了这套从数据准备、模型构建、训练调优到部署推理的完整流程后,你就具备了将图神经网络应用到真实项目中的基本能力。剩下的,就是在具体业务场景中不断打磨和深化了。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐