游乐游手机版
首页/AI热点日报/热点详情

详解如何使用PyTorch构建图神经网络从入门到精通

类型:热点整理2026-07-22
掌握图神经网络(GNN)的完整知识体系,包括GNN的定义、常见类型及其实际应用场景。同时,你将学习如何使用PyTorch框架从零构建图神经网络模型。 1 什么是图? 图由节点(Node)和边(Edge)构成——节点可以代表一个人、一个地点或一个物体,边则定义了节点之间的关联关系。边可以是单向的(有

掌握图神经网络(GNN)的完整知识体系,包括GNN的定义、常见类型及其实际应用场景。同时,你将学习如何使用PyTorch框架从零构建图神经网络模型。

1. 什么是图?

图由节点(Node)和边(Edge)构成——节点可以代表一个人、一个地点或一个物体,边则定义了节点之间的关联关系。边可以是单向的(有向)或双向的(无向),取决于依赖的方向性。简单来说,图中蓝色圆圈表示节点,箭头表示边,箭头方向指示了两个节点之间的依赖方向。

接下来看一个复杂图数据集的实例:爵士音乐家网络,该网络包含198个节点和2742条边。

爵士音乐家网络 https://datarepository.wolframcloud.com/resources/Jazz-Musicians-Network

在下面的社区图中,不同颜色的节点代表不同的音乐家群体,边表示他们之间的合作关系。这是一个协作网络——每位音乐家既与社区内部成员相连,也与社区外部成员相连。

图在解决涉及关系和相互作用的复杂问题时表现优异,广泛应用于模式识别、社交网络分析、推荐系统及语义分析等领域。构建基于图的解决方案正成为一个前沿方向,为复杂且相互关联的数据集提供深刻的洞察。

2. 使用 NetworkX 创建图

接下来,我们使用 NetworkX 库来创建图。本示例参考了 Daniel Holmberg 的博客《Python中的图神经网络》。步骤非常简单:首先创建一个 DiGraph 对象 'H',然后添加带有不同标签、颜色和大小的节点,再添加边来定义节点之间的依赖关系。例如 '(0,1)' 表示节点0对节点1有方向性依赖,再加入 '(1,0)' 则形成双向关系。最后,提取颜色和大小列表,调用 draw 函数绘制图形。

import networkx as nx

H = nx.DiGraph()

# adding nodes
H.add_nodes_from([
    (0, {"color": "blue", "size": 250}),
    (1, {"color": "yellow", "size": 400}),
    (2, {"color": "orange", "size": 150}),
    (3, {"color": "red", "size": 600})
])

# adding edges
H.add_edges_from([
    (0, 1),
    (1, 2),
    (1, 0),
    (1, 3),
    (2, 3),
    (3,0)
])

node_colors = nx.get_node_attributes(H, "color").values()
colors = list(node_colors)

node_sizes = nx.get_node_attributes(H, "size").values()
sizes = list(node_sizes)

# Plotting Graph
nx.draw(H, with_labels=True, node_color=colors, node_size=sizes)

下一步,用 to_undirected() 将图从有向图转换为无向图。

# 转换为无向图
G = H.to_undirected()
nx.draw(G, with_labels=True, node_color=colors, node_size=sizes)

3. 为什么分析图很难?

基于图的数据结构存在一些固有的挑战,数据分析师必须充分了解。图处于非欧几里得空间——它不在传统的2D或3D空间中,这使得数据解释更加困难。为了在2D平面上可视化,需要借助各种降维工具。图是动态的,没有固定形态。两个看起来完全不同的图,可能具有相似的邻接矩阵表示,这意味着传统统计工具难以直接应用。此外,随着规模和维度的增加,图会变得异常复杂——节点众多、边数成千上万,密集的拓扑结构使得理解和提取信息变得更加困难。

4. 什么是图神经网络(GNN)?

图神经网络(GNN)是一种专门用于处理图数据结构的特殊神经网络。它深受卷积神经网络(CNN)和图嵌入(Graph Embedding)的影响,广泛应用于节点预测、边预测以及基于图的任务。例如,CNN用于图像分类时处理的是像素网格,而GNN则作用于图结构;RNN用于文本分类时,GNN将每个单词视为图节点。GNN的诞生,正是为了解决传统CNN在面对任意大小和复杂结构的图数据时表现不佳的问题。

输入图经过多个神经网络层,被转换为图嵌入(Graph Embedding),从而保留节点、边以及全局上下文的信息。节点A和C的特征向量通过神经网络层进行聚合,然后传递到下一层。

4.1 图神经网络的类型

图神经网络目前有几种主流类型,大多都带有CNN的影子。下面介绍几种最流行的。

  • 图卷积网络(GCN):类似于传统CNN,通过检查相邻节点来学习特征。它聚合节点向量,传入全连接层,再通过激活函数引入非线性。简单来说,GCN = 图卷积 + 线性层 + 非线性激活。主要分为空间卷积和频谱卷积两类。
  • 图自编码器网络:使用编码器学习图的低维表示,再通过解码器重建输入图,中间由瓶颈层连接。常用于链路预测,因为自编码器在平衡类别数据方面表现出色。
  • 循环图神经网络(RGNN):学习最优扩散模式,能处理多关系图(即单个节点具有多种关系)。它通过正则化增强平滑性、消除过参数化,计算量小但效果显著。典型应用包括文本生成、机器翻译、语音识别、图像描述生成、视频标注和文本摘要。
  • 门控图神经网络(GGNN):在处理长期依赖任务上优于RGNN。它引入了节点门、边门和时间门,类似于门控循环单元(GRU),能够在不同状态下记住或遗忘信息。

4.2 图神经网络任务类型

下面列举一些常见的GNN任务类型及应用场景:

  • 图分类:将整个图划分到不同类别,常用于社交网络分析、文本分类等场景。
  • 节点分类:利用相邻节点的标签信息,预测图中缺失的节点标签。
  • 链路预测:预测邻接矩阵中缺失的一对节点之间是否存在链接,在社交网络分析中尤为常见。
  • 社区检测:根据边的结构将节点划分为不同簇,需从边权重、距离和图对象中学习。
  • 图嵌入:将图映射为低维向量,同时保留节点、边和结构的有效信息。
  • 图生成:从样本图的分布中学习,生成一个全新但结构相似的图。

图神经网络类型概览

4.3 图神经网络的缺点

使用GNN时也存在一些短板需要注意。大多数神经网络可以通过增加深度来提升性能,但GNN通常为浅层网络(一般不超过三层),这限制了其在大数据集上达到最先进性能的能力。此外,图结构不断变化,使得模型训练更加困难。部署到生产环境时还面临扩展性问题——由于GNN计算成本高,对于庞大复杂的图结构,在工业级应用中难以有效扩展。

5. 什么是图卷积网络(GCN)?

大多数GNN本质上都是图卷积网络(GCN),在进入节点分类实战之前,有必要先理解GCN。GCN中的卷积概念与CNN类似——将神经元与权重(滤波器)相乘,从数据特征中学习。在图像上,卷积如同滑动窗口,从相邻单元学习特征,滤波器通过权重共享来识别各种面部特征。而在图卷积中,模型从相邻节点学习特征。GCN与CNN的核心区别在于:GCN专门设计用于非欧几里得数据结构,节点和边的顺序可以变化。

CNN vs GCN

GCN有两种主要类型:

  • 空间图卷积网络:利用空间特征从位置空间中学习。
  • 频谱图卷积网络:利用图拉普拉斯矩阵的特征值分解进行节点间信息传播,灵感来源于信号与系统的波动传播。

6. 图神经网络如何工作?使用 PyTorch 构建图神经网络

接下来,我们将构建并训练一个用于节点分类的频谱图卷积模型。文末提供了完整代码,供你体验并运行第一个基于图的机器学习模型。

6.1 准备

首先安装PyTorch,因为 pytorch_geometric 依赖它。

!pip install -q torch

然后根据 PyTorch 版本安装 torch-scattertorch-sparse,最后从 GitHub 安装最新版 pytorch_geometric

%%capture
import os
import torch
os.environ['TORCH'] = torch.__version__
os.environ['PYTHONWARNINGS'] = "ignore"
!pip install torch-scatter -f https://data.pyg.org/whl/torch-${TORCH}.html
!pip install torch-sparse -f https://data.pyg.org/whl/torch-${TORCH}.html
!pip install git+https://github.com/pyg-team/pytorch_geometric.git

6.2 Planetoid Cora 数据集

Planetoid 数据集来自 Cora、CiteSeer 和 PubMed 三个引文网络。节点表示文档,每个节点具有1433维词袋(Bag-of-Words)特征向量;边表示论文之间的引用关系。共有7个类别,我们要训练模型来预测缺失的标签。导入数据后,对词袋输入进行行标准化,然后分析数据集及其第一个图对象。

from torch_geometric.datasets import Planetoid
from torch_geometric.transforms import NormalizeFeatures

dataset = Planetoid(root='data/Planetoid', name='Cora', transform=NormalizeFeatures())
print(f'Dataset: {dataset}:')
print('======================')
print(f'Number of graphs: {len(dataset)}')
print(f'Number of features: {dataset.num_features}')
print(f'Number of classes: {dataset.num_classes}')

data = dataset[0]  # Get the first graph object.
print(data)

Cora 数据集有2708个节点、10556条边、1433个特征、7个类别。第一个图对象包含训练、验证和测试掩码。

Dataset: Cora():
======================
Number of graphs: 1
Number of features: 1433
Number of classes: 7
Data(x=[2708, 1433], edge_index=[2, 10556], y=[2708], train_mask=[2708], val_mask=[2708], test_mask=[2708])

6.3 使用 GNN 进行节点分类

创建一个包含两个GCNConv层的GCN模型,使用ReLU激活函数和0.5的dropout率,隐藏通道数为16。

GCN 层的数学表达式:

其中 W(ℓ+1) 是可训练的权重矩阵,Cw,v 是每条边的固定标准化系数。

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

class GCN(torch.nn.Module):
    def __init__(self, hidden_channels):
        super().__init__()
        torch.manual_seed(1234567)
        self.conv1 = GCNConv(dataset.num_features, hidden_channels)
        self.conv2 = GCNConv(hidden_channels, dataset.num_classes)

    def forward(self, x, edge_index):
        x = self.conv1(x, edge_index)
        x = x.relu()
        x = F.dropout(x, p=0.5, training=self.training)
        x = self.conv2(x, edge_index)
        return x

model = GCN(hidden_channels=16)
print(model)
>>> GCN(
    (conv1): GCNConv(1433, 16)
    (conv2): GCNConv(16, 7)
  )

6.4 可视化未经训练的 GCN 网络

使用sklearn.manifold.TSNE和matplotlib可视化未训练网络的节点嵌入,将7维嵌入降维到2D散点图。

%matplotlib inline
import matplotlib.pyplot as plt
from sklearn.manifold import TSNE

def visualize(h, color):
    z = TSNE(n_components=2).fit_transform(h.detach().cpu().numpy())
    plt.figure(figsize=(10,10))
    plt.xticks([])
    plt.yticks([])
    plt.scatter(z[:, 0], z[:, 1], s=70, c=color, cmap="Set2")
    plt.show()

评估模型,将训练数据传入未训练模型,观察各类别节点分布:

model.eval()
out = model(data.x, data.edge_index)
visualize(out, color=data.y)

6.5 训练 GNN

使用Adam优化器和交叉熵损失函数训练100个epoch。训练函数步骤:清除梯度、前向传播、计算训练节点损失、反向传播、更新参数。测试函数:预测类别、提取最高概率标签、统计正确数量、计算准确率。

model = GCN(hidden_channels=16)
optimizer = torch.optim.Adam(model.parameters(), lr=0.01, weight_decay=5e-4)
criterion = torch.nn.CrossEntropyLoss()

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

def test():
    model.eval()
    out = model(data.x, data.edge_index)
    pred = out.argmax(dim=1)
    test_correct = pred[data.test_mask] == data.y[data.test_mask]
    test_acc = int(test_correct.sum()) / int(data.test_mask.sum())
    return test_acc

for epoch in range(1, 101):
    loss = train()
    print(f'Epoch: {epoch:03d}, Loss: {loss:.4f}')
GAT(
  (conv1): GATConv(1433, 8, heads=8)
  (conv2): GATConv(64, 7, heads=8)
)
.. .. .. ..
.. .. .. ..
Epoch: 098, Loss: 0.5989
Epoch: 099, Loss: 0.6021
Epoch: 100, Loss: 0.5799

6.6 模型评估

在未见过的测试集上评估,得到 81.5% 的准确率。

test_acc = test()
print(f'Test Accuracy: {test_acc:.4f}')
>>> 测试准确率:0.8150

可视化训练后模型的输出嵌入:

model.eval()
out = model(data.x, data.edge_index)
visualize(out, color=data.y)

可以看到,训练后的模型对同一类别的节点产生了更清晰的聚类。

6.7 训练 GATConv 模型

接下来,我们使用GATConv层替换GCNConv。图注意力网络(GAT)通过掩码自注意力机制克服了GCNConv的缺点,通常能取得更好的结果。你也可以尝试其他GNN层,并调优优化器、dropout率和隐藏通道数。下面设置第一层为8个注意力头(heads),第二层为1个头,dropout为0.6,隐藏通道数为8,学习率为0.005。修改测试函数以支持指定掩码(验证、测试),便于在训练过程中打印并记录分数。

from torch_geometric.nn import GATConv

class GAT(torch.nn.Module):
    def __init__(self, hidden_channels, heads):
        super().__init__()
        torch.manual_seed(1234567)
        self.conv1 = GATConv(dataset.num_features, hidden_channels,heads)
        self.conv2 = GATConv(heads*hidden_channels, dataset.num_classes,heads)

    def forward(self, x, edge_index):
        x = F.dropout(x, p=0.6, training=self.training)
        x = self.conv1(x, edge_index)
        x = F.elu(x)
        x = F.dropout(x, p=0.6, training=self.training)
        x = self.conv2(x, edge_index)
        return x

model = GAT(hidden_channels=8, heads=8)
print(model)
optimizer = torch.optim.Adam(model.parameters(), lr=0.005, weight_decay=5e-4)
criterion = torch.nn.CrossEntropyLoss()

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

def test(mask):
    model.eval()
    out = model(data.x, data.edge_index)
    pred = out.argmax(dim=1)
    correct = pred[mask] == data.y[mask]
    acc = int(correct.sum()) / int(mask.sum())
    return acc

val_acc_all = []
test_acc_all = []

for epoch in range(1, 101):
    loss = train()
    val_acc = test(data.val_mask)
    test_acc = test(data.test_mask)
    val_acc_all.append(val_acc)
    test_acc_all.append(test_acc)
    print(f'Epoch: {epoch:03d}, Loss: {loss:.4f}, Val: {val_acc:.4f}, Test: {test_acc:.4f}')
.. .. .. ..
.. .. .. ..
Epoch: 098, Loss: 1.1283, Val: 0.7960, Test: 0.8030
Epoch: 099, Loss: 1.1352, Val: 0.7940, Test: 0.8050
Epoch: 100, Loss: 1.1053, Val: 0.7960, Test: 0.8040

可以看出,本轮模型的表现并没有超越 GCNConv,还需要超参数优化或更多轮训练才能达到最佳。

6.8 模型评估

用折线图可视化验证和测试分数。

import numpy as np
plt.figure(figsize=(12,8))
plt.plot(np.arange(1, len(val_acc_all) + 1), val_acc_all, label='Validation accuracy', c='blue')
plt.plot(np.arange(1, len(test_acc_all) + 1), test_acc_all, label='Testing accuracy', c='red')
plt.xlabel('Epochs')
plt.ylabel('Accurarcy')
plt.title('GATConv')
plt.legend(loc='lower right', fontsize='x-large')
plt.sa vefig('gat_loss.png')
plt.show()

大约60轮后,验证和测试准确率稳定在 0.8 ± 0.02 左右。

再次可视化 GATConv 的节点聚类:

model.eval()
out = model(data.x, data.edge_index)
visualize(out, color=data.y)

GATConv 层在同类节点上也产生了同样的聚类效果。我们可以通过添加第二个验证集来减少过拟合,并尝试 pytorch_geometric 中多种不同的 GCN 层来提升模型性能。

GNN 常见问题

图神经网络(GNN)用于什么?

图神经网络(GNN)直接作用于图数据集,经过训练后可以预测节点、边以及图相关的各类任务。其具体应用包括图和节点分类、链路预测、图聚类与生成,以及图像和文本分类。

在图神经网络中,什么是图?

在图神经网络中,图是一种由节点(vertices)和连接边(edges)组成的数据结构。边可以是有向或无向的,具有动态形状和多维结构。例如在社交网络中,节点代表你的朋友,边代表你与每个人之间的关系。

图神经网络有多强大?

在图像和节点分类任务中,GNN通常优于传统CNN。许多GNN变体在节点和图分类领域已达到最先进水平(参见openreview.net)。

神经网络是否使用图论?

是的,神经网络与图论密切相关,尤其是那些专门处理非欧几里得数据的网络。有些神经网络本身即为图结构,或者其输出是图结构。

什么是图卷积网络?

图卷积网络(GCN)类似于专为图数据设计的卷积神经网络。它包含图卷积、线性层和非线性激活函数。GCN通过图上的滤波器检查节点和边,从而对图中的节点进行分类。

在深度学习中,什么是图?

图深度学习(Graph Deep Learning)也称为几何深度学习,通过堆叠多个神经网络层来提升性能。这是一个活跃的研究领域,科学家们正致力于在不牺牲性能的情况下加深网络层数。

来源:https://m.elecfans.com/article/2408190.html

相关热点

继续查看同栏目近期热点。

延伸阅读

补充最近整理过的热点入口。