PyG 图神经网络设计指南:从 MessagePassing 基类到异构图学习实战
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
本指南以 PyTorch Geometric(PyG)官方教程《Design of Graph Neural Networks》为主线,系统讲解如何利用MessagePassing基类从零实现 GCN、EdgeConv 等消息传递网络,并进一步掌握异构图(Heterogeneous Graph)的建模、变换与三类异构 GNN 构建方案。读完本文,你将能够自定义任意消息传递算子,并把同构图模型一键迁移到 ogbn-mag 等真实异构数据集上完成训练。
消息传递范式:GNN 的统一数学框架
把卷积算子推广到不规则域(如图、点云)上,通常被表达为**邻域聚合(neighborhood aggregation)或消息传递(message passing)**范式。设 $\mathbf{x}^{(k-1)}i \in \mathbb{R}^F$ 表示节点 $i$ 在第 $(k-1)$ 层的特征,$\mathbf{e}{j,i} \in \mathbb{R}^D$ 表示从节点 $j$ 指向节点 $i$ 的(可选)边特征,则消息传递图神经网络可以写成:
$$ \mathbf{x}_i^{(k)} = \gamma^{(k)} \left( \mathbf{x}i^{(k-1)}, \bigoplus{j \in \mathcal{N}(i)} , \phi^{(k)}\left(\mathbf{x}_i^{(k-1)}, \mathbf{x}j^{(k-1)},\mathbf{e}{j,i}\right) \right) $$
其中 $\bigoplus$ 表示一个可微的、置换不变的聚合函数,例如 sum、mean 或 max;$\gamma$ 与 $\phi$ 表示可微函数,例如 MLP(多层感知机)。直观地说,每个节点先从邻居收集"消息"(由 $\phi$ 生成),再按置换不变的方式聚合(由 $\bigoplus$ 完成),最后结合自身信息更新表示(由 $\gamma$ 完成)。几乎所有主流 GNN 算子——GCN、GraphSAGE、GAT、GIN、EdgeConv——都是该公式的特例,这正是 PyG 设计MessagePassing基类的理论根基。
MessagePassing 基类:PyG 提供的消息传递脚手架
PyG 在 message_passing.py 中提供了MessagePassing基类,它自动处理消息的传播(propagation)流程,用户只需定义三个要素:
- $\phi$,即
message()函数; - $\gamma$,即
update()函数; - 聚合方案,即
aggr="add"、aggr="mean"或aggr="max"。
构造函数参数
MessagePassing的构造签名(见 源码 L110-L118)如下:
| 参数 | 说明 | 默认值 |
|---|---|---|
aggr | 聚合方案:"add"/"sum"、"mean"、"min"、"max"、"mul",也可以是Aggregation模块或字符串列表(列表时各聚合结果在最后一维拼接) | "sum"(与"add"等价) |
aggr_kwargs | 传递给自动解析出的聚合函数的额外参数 | None |
flow | 消息传递方向:"source_to_target"或"target_to_source" | "source_to_target" |
node_dim | 沿哪个轴传播消息(批量维度处理) | -2 |
decomposed_layers | 特征分解层数,用于降低峰值内存、加速 CPU 推理 | 1 |
注意:源码中flow只接受'source_to_target'与'target_to_source'两个取值,传入其他值会抛出ValueError(见 L121-L123)。decomposed_layers通过把特征维切片成多个子层分批做聚合,可在 CPU 上显著降低峰值内存,但不适用于注意力类 GNN(消息无法简单分解)。
四个核心方法
propagate(edge_index, size=None, **kwargs):启动消息传播的入口(源码 L421)。它接收边索引以及构造消息、更新节点嵌入所需的全部附加数据。关键特性:propagate不只限于形状为 $[N, N]$ 的方阵邻接矩阵,也可以处理形状为 $[N, M]$ 的一般稀疏赋值矩阵(即二部图),只需传入size=(N, M);若为None则默认是方阵。对于具有两类独立节点的二部图,若两类节点各自持有信息,可用元组传入,例如x=(x_N, x_M)。message(...):对应于 $\phi$,为每条边 $(j,i) \in \mathcal{E}$(当flow="source_to_target"时)构造发往节点 $i$ 的消息。它可以接收任何传给propagate的参数;更妙的是,把变量名加上_i或_j后缀,PyG 会自动把张量映射到对应的目标/源节点,例如x_i、x_j。这里约定 $i$ 是聚合信息的中心节点,$j$ 是邻居节点。aggregate(...):执行 $\bigoplus$ 聚合(源码 L577)。默认委托给__init__中由aggr解析出的Aggregation模块。update(aggr_out, ...):对应于 $\gamma$,对每个节点 $i \in \mathcal{V}$ 更新节点嵌入(源码 L609)。第一个参数是聚合输出,其余参数来自propagate的输入。
propagate内部会依次调用message→aggregate→update(见 源码 L499-L550)。其底层实现中有两个值得关注的机制:
- Inspector 参数收集:构造时用
Inspector反射message/aggregate/update的签名,自动确定需要从propagate的**kwargs中收集哪些参数(L138-L151)。 - 张量抬升(lifting):
_lift通过index_select(self.node_dim, edge_index[dim])把节点特征按边索引展开成边级张量,这就是x_j的来源;_collect则根据flow决定_i/_j分别取edge_index的第 1/0 行。只要张量持有源或目标节点特征,任何张量都可以用_i/_j后缀自动抬升。
下面通过重实现两个经典算子——GCN 层(Kipf & Welling)与 EdgeConv 层(Wang et al.)——来验证这套机制。
实战一:从零实现 GCN 层
GCN 层的数学定义为:
$$ \mathbf{x}i^{(k)} = \sum{j \in \mathcal{N}(i) \cup { i }} \frac{1}{\sqrt{\deg(i)} \cdot \sqrt{\deg(j)}} \cdot \left( \mathbf{W}^{\top} \cdot \mathbf{x}_j^{(k-1)} \right) + \mathbf{b} $$
即:邻居节点特征先经权重矩阵 $\mathbf{W}$ 线性变换,再按度数归一化,最后求和;对聚合输出施加偏置向量 $\mathbf{b}$。这个公式可拆解为六个步骤:
- 给邻接矩阵添加自环(self-loops);
- 线性变换节点特征矩阵;
- 计算归一化系数;
- 在 $\phi$(
message)中归一化节点特征; - 求和聚合邻居节点特征(
"add"聚合); - 施加最终偏置向量。
步骤 1–3 通常在消息传递前完成,步骤 4–5 可直接交给MessagePassing基类。完整实现如下(与教程 create_gnn.rst 一致):
import torch from torch.nn import Linear, Parameter from torch_geometric.nn import MessagePassing from torch_geometric.utils import add_self_loops, degree class GCNConv(MessagePassing): def __init__(self, in_channels, out_channels): super().__init__(aggr='add') # "Add" aggregation (Step 5). self.lin = Linear(in_channels, out_channels, bias=False) self.bias = Parameter(torch.empty(out_channels)) self.reset_parameters() def reset_parameters(self): self.lin.reset_parameters() self.bias.data.zero_() def forward(self, x, edge_index): # x has shape [N, in_channels] # edge_index has shape [2, E] # Step 1: Add self-loops to the adjacency matrix. edge_index, _ = add_self_loops(edge_index, num_nodes=x.size(0)) # Step 2: Linearly transform node feature matrix. x = self.lin(x) # Step 3: Compute normalization. row, col = edge_index deg = degree(col, x.size(0), dtype=x.dtype) deg_inv_sqrt = deg.pow(-0.5) deg_inv_sqrt[deg_inv_sqrt == float('inf')] = 0 norm = deg_inv_sqrt[row] * deg_inv_sqrt[col] # Step 4-5: Start propagating messages. out = self.propagate(edge_index, x=x, norm=norm) # Step 6: Apply a final bias vector. out = out + self.bias return out def message(self, x_j, norm): # x_j has shape [E, out_channels] # Step 4: Normalize node features. return norm.view(-1, 1) * x_j逐段解读:
super().__init__(aggr='add')选择了"求和"聚合(步骤 5)。PyG 内部把'add'与'sum'视为同一聚合,见FUSE_AGGRS = {'add', 'sum', 'mean', 'min', 'max'}(message_passing.py L35)。add_self_loops来自 utils/loop.py,为edge_index补充 $(i,i)$ 自环(步骤 1)。PyG 官方实现的GCNConv(gcn_conv.py)更进一步使用add_remaining_self_loops避免重复添加,并通过gcn_norm(gcn_conv.py L45-L113)在稀疏矩阵与稠密张量两种表示下统一完成"自环 + 对称归一化"。degree(col, ...)来自 utils/_degree.py,计算每个节点的度数;deg_inv_sqrt中无穷大被置 0 以处理孤立节点(步骤 3)。propagate(edge_index, x=x, norm=norm)内部依次调用message、aggregate、update。在message(self, x_j, norm)中,x_j是抬升(lifted)张量,包含每条边源节点的特征(即各节点的邻居特征),norm同样按边索引展开后与x_j逐元素相乘完成归一化(步骤 4)。
初始化与调用极其简单,该层可直接作为深度架构的积木:
conv = GCNConv(16, 32) x = conv(x, edge_index)实战二:实现 Edge 卷积层
EdgeConv 层用于处理图或点云,数学定义为:
$$ \mathbf{x}i^{(k)} = \max{j \in \mathcal{N}(i)} h_{\mathbf{\Theta}} \left( \mathbf{x}_i^{(k-1)}, \mathbf{x}_j^{(k-1)} - \mathbf{x}_i^{(k-1)} \right) $$
其中 $h_{\mathbf{\Theta}}$ 是一个 MLP。与 GCN 类似,这次改用"max"聚合:
import torch from torch.nn import Sequential as Seq, Linear, ReLU from torch_geometric.nn import MessagePassing class EdgeConv(MessagePassing): def __init__(self, in_channels, out_channels): super().__init__(aggr='max') # "Max" aggregation. self.mlp = Seq(Linear(2 * in_channels, out_channels), ReLU(), Linear(out_channels, out_channels)) def forward(self, x, edge_index): # x has shape [N, in_channels] # edge_index has shape [2, E] return self.propagate(edge_index, x=x) def message(self, x_i, x_j): # x_i has shape [E, in_channels] # x_j has shape [E, in_channels] tmp = torch.cat([x_i, x_j - x_i], dim=1) # tmp has shape [E, 2 * in_channels] return self.mlp(tmp)在message中,我们同时用到了x_i与x_j:x_i是目标节点特征,x_j - x_i是相对源节点特征(刻画了每条边两端点的"差分"信息),二者在特征维拼接后送入 MLP。官方 edge_conv.py 中的EdgeConv实现与此一致(L60-L61),差别仅在于官方把 MLP 作为nn参数从外部注入、message用dim=-1拼接,且forward支持二部图输入(x = (x, x))。
EdgeConv 本质上是动态卷积:每一层都在特征空间中用最近邻重新构造图。PyG 提供了 GPU 加速的批量 k-NN 图生成方法torch_geometric.nn.pool.knn_graph:
from torch_geometric.nn import knn_graph class DynamicEdgeConv(EdgeConv): def __init__(self, in_channels, out_channels, k=6): super().__init__(in_channels, out_channels) self.k = k def forward(self, x, batch=None): edge_index = knn_graph(x, self.k, batch, loop=False, flow=self.flow) return super().forward(x, edge_index)knn_graph计算最近邻图,再调用EdgeConv.forward完成消息传递。官方 DynamicEdgeConv 还额外要求pyg-lib>=0.6.0支持,并支持num_workers并行计算 k-NN。调用接口干净利落:
conv = DynamicEdgeConv(3, 128, k=6) x = conv(x, batch)动手练习:验证你的理解
教程 create_gnn.rst 提供了一个练习数据集:
import torch from torch_geometric.data import Data edge_index = torch.tensor([[0, 1], [1, 0], [1, 2], [2, 1]], dtype=torch.long) x = torch.tensor([[-1], [0], [1]], dtype=torch.float) data = Data(x=x, edge_index=edge_index.t().contiguous())围绕GCNConv思考以下问题(答案都可从源码与上述解读推出):
row和col分别保存什么信息?(row = edge_index[0]为源节点,col = edge_index[1]为目标节点)degree函数做了什么?(统计每个节点的度数,即入边数量)- 为什么用
degree(col, ...)而不是degree(row, ...)?(flow="source_to_target"下目标节点按col聚合消息,度数应统计目标端) deg_inv_sqrt[col]和deg_inv_sqrt[row]分别做什么?(对称归一化中目标端与源端的归一化因子)- 在
message中x_j保存什么?若self.lin为恒等函数,x_j的具体内容是什么?(每条边源节点的原始特征) - 给
GCNConv添加一个update函数,把经变换的中心节点特征加到聚合输出上(相当于实现残差连接)。
围绕EdgeConv的问题:
x_i和x_j - x_i分别是什么?(目标节点特征;边两端点特征之差)torch.cat([x_i, x_j - x_i], dim=1)做了什么?为什么是dim=1?(按特征维拼接,因为每行是一条边、每列是一个特征通道)
异构图学习:为什么需要专门的数据结构
大量真实数据集以异构图(heterogeneous graph)形式存储——推荐领域的社交图就是典型例子:图中同时存在多种实体类型(用户、商品、商家)与多种关系类型。这类图上,不同节点/边类型携带不同维度、不同类型的特征,单一的特征张量无法容纳整张图的信息,因此需要按类型分别维护数据张量;相应地,消息传递公式也需允许消息函数与更新函数按节点/边类型条件化。
引导示例:ogbn-mag 网络
教程以 OGB 数据集套件中的ogbn-mag网络为引导示例(示意图见 hg_example.svg):
该图共有1,939,743 个节点,分为四种节点类型:author(作者)、paper(论文)、institution(机构)与field of study(研究领域);共有21,111,007 条边,也分四种类型:
- writes:作者撰写某篇论文;
- affiliated with:作者隶属于某机构;
- cites:论文引用另一篇论文;
- has topic:论文属于某个研究领域。
任务目标是:根据图中存储的信息推断每篇论文的发表 venue(会议或期刊)。
创建异构图:HeteroData 详解
首先创建torch_geometric.data.HeteroData对象,按类型分别定义节点特征张量、边索引张量与边特征张量:
from torch_geometric.data import HeteroData data = HeteroData() data['paper'].x = ... # [num_papers, num_features_paper] data['author'].x = ... # [num_authors, num_features_author] data['institution'].x = ... # [num_institutions, num_features_institution] data['field_of_study'].x = ... # [num_field, num_features_field] data['paper', 'cites', 'paper'].edge_index = ... # [2, num_edges_cites] data['author', 'writes', 'paper'].edge_index = ... # [2, num_edges_writes] data['author', 'affiliated_with', 'institution'].edge_index = ... # [2, num_edges_affiliated] data['paper', 'has_topic', 'field_of_study'].edge_index = ... # [2, num_edges_topic] data['paper', 'cites', 'paper'].edge_attr = ... # [num_edges_cites, num_features_cites] data['author', 'writes', 'paper'].edge_attr = ... # [num_edges_writes, num_features_writes] data['author', 'affiliated_with', 'institution'].edge_attr = ... # [num_edges_affiliated, num_features_affiliated] data['paper', 'has_topic', 'field_of_study'].edge_attr = ... # [num_edges_topic, num_features_topic]节点/边张量在首次访问时自动创建,并以字符串键索引。节点类型用单个字符串标识;边类型用三元组(source_node_type, edge_type, destination_node_type)标识。因此该数据对象天然允许每种类型拥有不同的特征维度。
按属性名(而非按节点/边类型)分组后的异构字典可直接作为 GNN 模型的输入:
model = HeteroGNN(...) output = model(data.x_dict, data.edge_index_dict, data.edge_attr_dict)直接加载 OGB_MAG 数据集
若数据集在 PyG 的数据集列表中,可直接导入使用——数据集会被自动下载到root并完成预处理:
from torch_geometric.datasets import OGB_MAG dataset = OGB_MAG(root='./data', preprocess='metapath2vec') data = dataset[0]打印该data对象验证结构:
HeteroData( paper={ x=[736389, 128], y=[736389], train_mask=[736389], val_mask=[736389], test_mask=[736389] }, author={ x=[1134649, 128] }, institution={ x=[8740, 128] }, field_of_study={ x=[59965, 128] }, (author, affiliated_with, institution)={ edge_index=[2, 1043998] }, (author, writes, paper)={ edge_index=[2, 7145660] }, (paper, cites, paper)={ edge_index=[2, 5416271] }, (paper, has_topic, field_of_study)={ edge_index=[2, 7505078] } )注意:原始ogbn-mag网络只给 "paper" 节点提供特征。PyG 的OGB_MAG提供了下载已处理版本的选项——用"metapath2vec"或"TransE"得到的结构特征填充到无特征节点上,这正是 OGB 排行榜顶部方案常用的做法(见 datasets/ogb_mag.py)。
实用工具函数
HeteroData提供多种修改与分析图的工具:
# 单独索引某个节点或边存储 paper_node_data = data['paper'] cites_edge_data = data['paper', 'cites', 'paper'] # 若边类型可由节点对或边类型唯一确定,可简写 cites_edge_data = data['paper', 'paper'] cites_edge_data = data['cites'] # 新增/删除节点类型或张量 data['paper'].year = ... # Setting a new paper attribute del data['field_of_study'] # Deleting 'field_of_study' node type del data['has_topic'] # Deleting 'has_topic' edge type # 查看元数据(所有节点/边类型) node_types, edge_types = data.metadata() print(node_types) ['paper', 'author', 'institution'] print(edge_types) [('paper', 'cites', 'paper'), ('author', 'writes', 'paper'), ('author', 'affiliated_with', 'institution')] # 设备迁移 data = data.to('cuda:0') data = data.cpu() # 图性质分析 data.has_isolated_nodes() data.has_self_loops() data.is_undirected()还可通过to_homogeneous()转换为同构"带类型"图(当各类型特征维度一致时可保留特征):
homogeneous_data = data.to_homogeneous() print(homogeneous_data) Data(x=[1879778, 128], edge_index=[2, 13605929], edge_type=[13605929])其中homogeneous_data.edge_type是一个边级向量,以整数记录每条边的边类型。
异构图变换
大多数预处理变换同样适用于异构data对象:
import torch_geometric.transforms as T data = T.ToUndirected()(data) data = T.AddSelfLoops()(data) data = T.NormalizeFeatures()(data)ToUndirected:把有向图转换为(PyG 表示下的)无向图——为所有边添加反向边,必要时会为异构图添加反向边类型,使后续消息传递沿两个方向进行;AddSelfLoops:对类型为'node_type'的所有节点、以及所有形如('node_type', 'edge_type', 'node_type')的现有边类型添加自环,结果是每个节点可能收到一个或多个(每个合适的边类型一个)来自自身的信息;NormalizeFeatures:与同构图一致,把所有类型的指定特征归一化到和为 1。
创建异构 GNN 的三种方式
标准消息传递 GNN(MP-GNN)无法直接应用于异构数据:不同类型的节点/边特征不能用同一函数处理(维度、语义不同)。自然的思路是为每种边类型单独实现消息函数、每种节点类型单独实现更新函数——运行时按边类型字典迭代计算消息、按节点类型字典更新节点。为避免不必要的运行时开销并简化建模,PyG 提供了三种构建异构 GNN 模型的方式:
- 自动转换:用
torch_geometric.nn.to_hetero或to_hetero_with_bases把同构 GNN 自动转换为异构 GNN; - HeteroConv 包装器:用
torch_geometric.nn.conv.HeteroConv为不同边类型定义各自的卷积; - 直接部署现成(或自研)异构算子:如
torch_geometric.nn.conv.HGTConv。
方式一:自动转换同构 GNN(to_hetero / to_hetero_with_bases)
PyG 内置to_hetero与to_hetero_with_bases函数(实现见 to_hetero_transformer.py),可把任意 PyG GNN 模型自动转换为异构输入模型。以 to_hetero_mag.py 为例:
import torch_geometric.transforms as T from torch_geometric.datasets import OGB_MAG from torch_geometric.nn import SAGEConv, to_hetero dataset = OGB_MAG(root='./data', preprocess='metapath2vec', transform=T.ToUndirected()) data = dataset[0] class GNN(torch.nn.Module): def __init__(self, hidden_channels, out_channels): super().__init__() self.conv1 = SAGEConv((-1, -1), hidden_channels) self.conv2 = SAGEConv((-1, -1), out_channels) def forward(self, x, edge_index): x = self.conv1(x, edge_index).relu() x = self.conv2(x, edge_index) return x model = GNN(hidden_channels=64, out_channels=dataset.num_classes) model = to_hetero(model, data.metadata(), aggr='sum')转换过程会复制消息函数,使其按每种边类型独立工作,如下图所示(来源 to_hetero.svg)。转换后,模型期望的输入从同构图中的单一张量变为以节点/边类型为键的字典。注意SAGEConv传入的是(-1, -1)形状的in_channels元组——这是为了支持二部图(异构边类型)的消息传递。
惰性初始化(lazy initialization):由于不同类型间的输入特征数与张量尺寸各不相同,PyG 用-1作为in_channels表示惰性初始化,避免手动计算计算图中所有张量尺寸;惰性初始化对所有 PyG 算子生效。只需调用一次模型即可完成参数初始化:
with torch.no_grad(): # Initialize lazy modules. out = model(data.x_dict, data.edge_index_dict)to_hetero/to_hetero_with_bases对可自动转换的同构架构非常灵活,跳连(skip-connection)、Jumping Knowledge 等技术开箱即用。例如,实现带可学习跳连的异构图注意力网络只需:
from torch_geometric.nn import GATConv, Linear, to_hetero class GAT(torch.nn.Module): def __init__(self, hidden_channels, out_channels): super().__init__() self.conv1 = GATConv((-1, -1), hidden_channels, add_self_loops=False) self.lin1 = Linear(-1, hidden_channels) self.conv2 = GATConv((-1, -1), out_channels, add_self_loops=False) self.lin2 = Linear(-1, out_channels) def forward(self, x, edge_index): x = self.conv1(x, edge_index) + self.lin1(x) x = x.relu() x = self.conv2(x, edge_index) + self.lin2(x) return x model = GAT(hidden_channels=64, out_channels=dataset.num_classes) model = to_hetero(model, data.metadata(), aggr='sum')这里特意用add_self_loops=False关闭自环:在二部图中"自环"概念不成立(边类型两端节点类型不同),若开启会错误地在二部图上添加[(0, 0), (1, 1), ...]这类边。为保留中心节点信息,改用可学习跳连conv(x, edge_index) + lin(x):注意力消息从源节点传到目标节点,输出再与既有目标节点特征相加。
转换后的模型按标准流程训练:
def train(): model.train() optimizer.zero_grad() out = model(data.x_dict, data.edge_index_dict) mask = data['paper'].train_mask loss = F.cross_entropy(out['paper'][mask], data['paper'].y[mask]) loss.backward() optimizer.step() return float(loss)方式二:使用异构卷积包装器 HeteroConv
HeteroConv(hetero_conv.py)允许从零为每种边类型定义自定义消息与更新函数,构建任意异构 MP-GNN。与to_hetero(所有边类型共用同一算子)不同,包装器为不同边类型指定不同算子:它接收一个以边类型为键的子模块字典。以 hetero_conv_dblp.py 为例:
import torch_geometric.transforms as T from torch_geometric.datasets import OGB_MAG from torch_geometric.nn import HeteroConv, GCNConv, SAGEConv, GATConv, Linear dataset = OGB_MAG(root='./data', preprocess='metapath2vec', transform=T.ToUndirected()) data = dataset[0] class HeteroGNN(torch.nn.Module): def __init__(self, hidden_channels, out_channels, num_layers): super().__init__() self.convs = torch.nn.ModuleList() for _ in range(num_layers): conv = HeteroConv({ ('paper', 'cites', 'paper'): GCNConv(-1, hidden_channels), ('author', 'writes', 'paper'): SAGEConv((-1, -1), hidden_channels), ('paper', 'rev_writes', 'author'): GATConv((-1, -1), hidden_channels, add_self_loops=False), }, aggr='sum') self.convs.append(conv) self.lin = Linear(hidden_channels, out_channels) def forward(self, x_dict, edge_index_dict): for conv in self.convs: x_dict = conv(x_dict, edge_index_dict) x_dict = {key: x.relu() for key, x in x_dict.items()} return self.lin(x_dict['author']) model = HeteroGNN(hidden_channels=64, out_channels=dataset.num_classes, num_layers=2)初始化与训练方式同前:
with torch.no_grad(): # Initialize lazy modules. out = model(data.x_dict, data.edge_index_dict)从源码看,HeteroConv内部会为字典中每个边类型检查add_self_loops设置,并警告那些"只作为源类型出现、从不作为目标类型更新表示"的节点类型(hetero_conv.py L70-L80);指向同一目标节点的多种关系的结果会按aggr聚合("sum"/"mean"/"min"/"max"/"cat"/None,见 L13-L26)。
方式三:部署现成的异构算子(如 HGTConv)
PyG 提供专为异构图设计的算子,例如torch_geometric.nn.conv.HGTConv,可直接用于搭建异构 GNN 模型(参考 hgt_dblp.py):
import torch_geometric.transforms as T from torch_geometric.datasets import OGB_MAG from torch_geometric.nn import HGTConv, Linear dataset = OGB_MAG(root='./data', preprocess='metapath2vec', transform=T.ToUndirected()) data = dataset[0] class HGT(torch.nn.Module): def __init__(self, hidden_channels, out_channels, num_heads, num_layers): super().__init__() self.lin_dict = torch.nn.ModuleDict() for node_type in data.node_types: self.lin_dict[node_type] = Linear(-1, hidden_channels) self.convs = torch.nn.ModuleList() for _ in range(num_layers): conv = HGTConv(hidden_channels, hidden_channels, data.metadata(), num_heads, group='sum') self.convs.append(conv) self.lin = Linear(hidden_channels, out_channels) def forward(self, x_dict, edge_index_dict): for node_type, x in x_dict.items(): x_dict[node_type] = self.lin_dictnode_type.relu_() for conv in self.convs: x_dict = conv(x_dict, edge_index_dict) return self.lin(x_dict['author']) model = HGT(hidden_channels=64, out_channels=dataset.num_classes, num_heads=2, num_layers=2)同样地,初始化与训练流程与前面完全一致(先with torch.no_grad()调用一次触发惰性初始化,再按标准训练函数迭代)。
大规模异构图:采样器与 mini-batch 训练
PyG 为异构图采样提供了多种功能:标准torch_geometric.loader.NeighborLoader同时支持同构与异构图,也有专用异构采样器如torch_geometric.loader.HGTLoader。这对大规模异构图的高效表示学习尤为重要——全量处理邻居在计算上不可承受。所有异构图加载器输出一个HeteroData对象(原数据的子集),差异主要在采样过程,因此从全批量训练切换到 mini-batch 训练只需极少的代码改动。
用NeighborLoader做邻居采样的示例(同样参考 to_hetero_mag.py):
import torch_geometric.transforms as T from torch_geometric.datasets import OGB_MAG from torch_geometric.loader import NeighborLoader transform = T.ToUndirected() # Add reverse edge types. data = OGB_MAG(root='./data', preprocess='metapath2vec', transform=transform)[0] train_loader = NeighborLoader( data, # Sample 15 neighbors for each node and each edge type for 2 iterations: num_neighbors=[15] * 2, # Use a batch size of 128 for sampling training nodes of type "paper": batch_size=128, input_nodes=('paper', data['paper'].train_mask), ) batch = next(iter(train_loader))NeighborLoader在异构图上还可按边类型做更细粒度的采样控制(非必需):
num_neighbors = {key: [15] * 2 for key in data.edge_types}input_nodes参数指定采样的局部邻域起点类型与索引——这里即train_mask标记的全部 "paper" 训练节点。打印batch得到:
HeteroData( paper={ x=[20799, 256], y=[20799], train_mask=[20799], val_mask=[20799], test_mask=[20799], batch_size=128 }, author={ x=[4419, 128] }, institution={ x=[302, 128] }, field_of_study={ x=[2605, 128] }, (author, affiliated_with, institution)={ edge_index=[2, 0] }, (author, writes, paper)={ edge_index=[2, 5927] }, (paper, cites, paper)={ edge_index=[2, 11829] }, (paper, has_topic, field_of_study)={ edge_index=[2, 10573] }, (institution, rev_affiliated_with, author)={ edge_index=[2, 829] }, (paper, rev_writes, author)={ edge_index=[2, 5512] }, (field_of_study, rev_has_topic, paper)={ edge_index=[2, 10499] } )batch共包含 28,187 个节点,用于计算 128 个 "paper" 节点的嵌入。采样节点总是按采样顺序排序,因此batch['paper']中前batch['paper'].batch_size个节点即原始 mini-batch 节点集合,通过切片即可方便地取出最终输出嵌入。
mini-batch 训练与全批量训练类似,区别在于遍历train_loader产生的 mini-batch 并逐个优化:
def train(): model.train() total_examples = total_loss = 0 for batch in train_loader: optimizer.zero_grad() batch = batch.to('cuda:0') batch_size = batch['paper'].batch_size out = model(batch.x_dict, batch.edge_index_dict) loss = F.cross_entropy(out['paper'][:batch_size], batch['paper'].y[:batch_size]) loss.backward() optimizer.step() total_examples += batch_size total_loss += float(loss) * batch_size return total_loss / total_examples关键点:损失计算只使用前 128 个 "paper" 节点——通过batch['paper'].batch_size对标签batch['paper'].y与预测out['paper']同时切片,二者分别对应原始 mini-batch 节点的标签与最终输出。
进一步阅读
- 本指南对应的原始教程:gnn_design.rst(gallery 入口)、create_gnn.rst、heterogeneous.rst
- 基类与算子源码:message_passing.py、gcn_conv.py、edge_conv.py、hetero_conv.py、to_hetero_transformer.py
- 完整可运行示例:to_hetero_mag.py、hetero_conv_dblp.py、hgt_dblp.py、bipartite_sage.py
- 异构数据与数据集:hetero_data.py、ogb_mag.py
- 异构图测试覆盖:test/nn/conv、test/data/test_hetero_data.py
掌握MessagePassing基类与三种异构建模方案后,你既可以按需定制任意图算子,也能把成熟同构模型快速迁移到真实的异构大规模图数据上,是进阶 PyG 深度定制与工业级落地的关键一步。
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考