首页 > 教程攻略 > ai资讯 >图算法之GraphSAGE原理以及代码实现

图算法之GraphSAGE原理以及代码实现

来源:互联网 时间:2026-08-04 14:23:33

一、GraphSAGE实现原理

图算法之GraphSAGE原理以及代码实现

先说几个核心判断。在Graph Embedding领域,早期的方法像DeepWalk、Node2Vec甚至GCN,都属于直推式(transductive)模式——它们训练时看见整个图,每个节点对应一个固定的Embedding向量,一旦图结构发生变化(比如新增节点),整张图就得重新训练。而GraphSAGE走的是另一条路:归纳式(inductive)。它不学节点的静态向量,而是学一个“如何从局部邻居信息生成节点嵌入”的映射函数。换句话说,拿到一个新节点,只要知道它的邻居长什么样,就能现场计算出它的Embedding。这意味着模型天然支持大规模图数据,也能轻松应对动态变化的图结构。

直推式学习与归纳式学习:

  1. 直推式学习(Transductive Learning):

    直推式模式下,模型只在训练集里出现过的样本上做预测,对未见过的样本无能为力。它适合的场景是:你手头只有一批固定节点,只需要给它们打上标签,不需要泛化到新数据。

  2. 归纳式学习(Inductive Learning):

    归纳式模式的目标是从训练数据中提取可迁移的规律,让模型能对从未见过的新样本做出合理预测。这才是真正意义上的“泛化”,也是生产环境中最常用的方式。

回到GraphSAGE本身。SAGE全称是SAmple and aggreGatE,核心动作就是两步——采样邻居,聚合信息。它不依赖全局图结构,而是为每个节点定义一个局部邻居采样范围,然后训练一组聚合函数(aggregator functions),这些函数能从不同跳数(hops)的邻居那里不断提取特征。推理时,就算来了一个完全陌生的节点,只要按同样的方式采样邻居并调用聚合函数,就能输出它的Embedding。

GraphSAGE实现步骤:

  1. a. 采样邻居

    :对每个目标节点,从它的邻居中随机抽取固定数量的节点。这么做是为了控制计算量,让算法能扩展到百万级节点。
  2. b. 聚合邻居信息

    :定义多种聚合函数(均值、池化、LSTM等),把采样到的邻居节点特征整合成一条统一的表示。
  3. c. 更新节点嵌入

    :将聚合后的邻居特征与目标节点自身的特征拼接(或求和),再通过一个全连接层做非线性变换,得到当前层的嵌入。
  4. d. 重复与优化

    :对整个图的所有节点迭代上述过程,通过反向传播训练聚合函数和变换层的参数。

上图中展示了为红色目标节点生成Embedding的流程。k表示搜索深度:k=1时采样3个直接邻居,k=2时采样5个二跳邻居。具体来说:第一步采样邻居节点;第二步将邻居信息逐层聚合,更新目标节点嵌入;第三步用这个嵌入去完成下游预测任务。

二、GraphSAGE伪代码

这里的K对应网络的层数,也决定了每个节点能聚合的跳数。例如K=2时,节点可以聚合它两跳邻居的信息。每一层循环中,先对节点v的所有邻居,用上一层的Embedding聚合出邻居的当前层表示,然后和v的上一层表示拼接,过非线性变换,得到v的当前层Embedding。

三、GraphSAGE的聚合器

聚合函数要解决的问题是:把一个无序的向量集合压缩成一个向量。图上的邻居没有天然顺序,所以聚合器必须是对称的——无论邻居节点以什么顺序输入,输出结果都一样。同时,还得有足够的表达能力。

  • Mean aggregator

    :最直接的方法——把目标节点和邻居节点的上一层向量拼接起来,然后对每个维度取均值。简单、稳定、快速。
  • LSTM aggregator

    :表达能力更强,但LSTM本身对输入顺序敏感。为了绕开这个限制,实践中先把邻居节点的向量集合随机打乱,再喂给LSTM,相当于用随机化消除顺序影响。
  • Pooling aggregator

    :让所有邻居节点共享一个权重矩阵,先过非线性全连接层,再对每个维度做max-pooling(或mean-pooling)。效果通常不错,且计算效率高。

四、GraphSAGE的损失函数

有监督损失函数

:直接根据下游任务来定。如果是节点分类,就用常规的交叉熵损失。

无监督损失函数

损失函数分两部分。蓝色部分希望:如果节点u和v在图上是相邻或接近的,那么它们的嵌入向量内积应该很大,经过Sigmoid后接近1,log损失趋近于0。粉色部分则做反向约束:如果u和v在图上距离很远,它们的嵌入内积应该是一个绝对值很大的负数(相差超过90度),经过Sigmoid后同样接近1,损失也为0。实际训练时,远离节点v的“负样本”u可能远多于正样本,所以从远距离分布中随机采样一部分负节点,再添加一个很小的epsilon防止log(0)。这样正负样本平衡,模型才能学到有意义的嵌入。

五、GraphSAGE代码实现

import torch
import torch.nn as nn
import torch.nn.functional as F
from torch_geometric.nn import SAGEConv
from torch_geometric.datasets import Planetoid
from torch_geometric.data import DataLoader

# 导入所需的库
# torch:PyTorch核心库
# torch.nn:包含了神经网络层的库
# torch.nn.functional:包含了神经网络函数的库
# torch_geometric.nn:PyTorch Geometric中的图神经网络模块
# torch_geometric.datasets:包含了一些常用的图数据集
# torch_geometric.data:定义了用于处理图数据的数据结构和函数

class GraphSage(nn.Module):
    def __init__(self, in_channels, hidden_channels, num_layers):
        super(GraphSage, self).__init__()
        self.convs = nn.ModuleList()
        self.convs.append(SAGEConv(in_channels, hidden_channels))
        for _ in range(num_layers - 1):
            self.convs.append(SAGEConv(hidden_channels, hidden_channels))

    def forward(self, x, edge_index):
        for conv in self.convs:
            x = conv(x, edge_index)
            x = F.relu(x)
        return x

# 定义GraphSage模型类,继承自nn.Module
# 初始化方法__init__中定义了GraphSage模型的网络层
# forward方法中定义了模型的前向传播过程

# 加载数据集
dataset = Planetoid(root='/tmp/Cora', name='Cora')

# 使用Planetoid加载Cora数据集,存储在/tmp/Cora目录下

# 划分数据集
train_loader = DataLoader(dataset, batch_size=64, shuffle=True)

# 将数据集划分为mini-batch,每个batch大小为64,进行随机打乱

# 实例化GraphSage模型
model = GraphSage(in_channels=dataset.num_features, hidden_channels=16, num_layers=2)

# 创建GraphSage模型的实例,指定输入特征维度、隐藏层维度和层数

# 训练循环
optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
criterion = nn.CrossEntropyLoss()

# 使用Adam优化器和交叉熵损失函数进行模型训练

def train():
    model.train()
    total_loss = 0
    for data in train_loader:
        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()
        total_loss += loss.item()
    return total_loss / len(train_loader)

# 定义训练函数train,对模型进行训练
# 遍历每个mini-batch,计算损失并更新模型参数

# 训练模型
for epoch in range(100):
    loss = train()
    print(f'Epoch {epoch + 1}, Loss: {loss:.4f}')

# 训练模型100个epoch,打印每个epoch的损失值