PyTorch Geometric 超图卷积实战指南:一篇讲透"群聊式"高阶关系建模
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
PyTorch Geometric(PyG)除了常规"一条边连两个节点"的图数据,还提供了超图卷积层。本文用HyperGraphData和HypergraphConv两个模块,讲清楚如何把群聊成员、多原子键合、"用户-商品-标签"这类高阶关系建模成超边,并训练一个节点分类器。
🧩 第一步:建一个"两个群聊"的超图数据
假设 5 个用户构成两个群聊:群 0 由用户 0、1、2 组成,群 1 由用户 1、2、3、4 组成。普通图里你得在群成员之间两两补边,而超图里每个群聊本身就是一条超边。
PyG 没有为超边单独发明张量格式,而是复用普通edge_index的布局:第一行存节点索引,第二行存该节点所属的超边编号。下面这段代码把上面两个群聊写成数据对象:
import torch from torch_geometric.data.hypergraph_data import HyperGraphData x = torch.randn(5, 16) # 5 个用户,每个 16 维特征 y = torch.tensor([0, 0, 1, 1, 2]) # 节点分类标签 edge_index = torch.tensor([ [0, 1, 2, 1, 2, 3, 4], # 第一行:节点索引 [0, 0, 0, 1, 1, 1, 1], # 第二行:节点属于哪条超边(群聊) ]) data = HyperGraphData(x=x, edge_index=edge_index, y=y) print(data.num_nodes, data.num_edges) # 5, 2这里有两个容易踩的坑:
- 构造参数名是
edge_index,而不是hyperedge_index。类定义在torch_geometric/data/hypergraph_data.py,但它没有挂到torch_geometric.data的顶层导出上,需要按上面那样从子模块直接 import。 num_edges按"第二行最大值 + 1"推断,所以超边编号必须从 0 开始连续;num_nodes则由第一行最大值推断。
🧮 超图卷积:一次"出去再回来"的两段聚合
HypergraphConv来自论文 "Hypergraph Convolution and Hypergraph Attention"(实现位于torch_geometric/nn/conv/hypergraph_conv.py)。把公式先翻译成人话:每个节点的特征先被聚合进它所在的每条超边,超边再把聚合结果送回每个节点,每一段都除以对应的"度"做归一化。对应公式为:
$$\mathbf{X}' = \mathbf{D}^{-1}\mathbf{H}\mathbf{W}\mathbf{B}^{-1}\mathbf{H}^{\top}\mathbf{X}\boldsymbol{\Theta}$$
其中 H 是 0/1 关联矩阵(被压缩成 edge_index 这种稀疏形式),W 是超边权重,D、B 分别是节点所属超边数的倒数和超边规模的倒数,Θ 是可学习的线性权重。两段聚合在源码forward中体现为对同一个hyperedge_index做两次propagate。
最小调用方式和普通 GCN 一致:
from torch_geometric.nn import HypergraphConv conv = HypergraphConv(16, 32) x = conv(x, edge_index) # 输出仍是每节点 32 维特征forward还接受两个可选参数:hyperedge_weight(长度为超边数 M 的向量,控制每条超边的权重,缺省为全 1)和num_edges(一般不用传,可从索引推断)。
⚖️ 注意力什么时候开,node 和 edge 模式怎么选
设置use_attention=True后,PyG 会为每个"节点—超边"关联打分,此时必须同时传入hyperedge_attr(源码里有显式 assert):形状为 [M, F] 的张量,表示每条超边自身的特征,例如群聊成员的平均画像或标签词统计。
edge_attr = torch.randn(2, 16) # 每个群聊一条 16 维特征 conv = HypergraphConv(16, 32, use_attention=True, heads=2, concat=True) x = conv(x, edge_index, hyperedge_attr=edge_attr)attention_mode决定 softmax 沿哪个维度归一化,两种模式回答的问题不同:
node(默认):在同一条超边内、所有属于它的节点之间算注意力,回答"这个群里谁贡献大";edge:对同一个节点、跨它所属的所有超边算注意力,回答"这个用户更关注哪个群"。
多头部分,concat=False时各头输出取平均而不是拼接,输出维度是out_channels而不是heads * out_channels。
✅ 串起来:一个两层超图分类器
把前面拼成完整模型,注意注意力模式下每层都要传edge_attr:
import torch.nn.functional as F class HyperGNN(torch.nn.Module): def __init__(self): super().__init__() self.conv1 = HypergraphConv(16, 32, use_attention=True) self.conv2 = HypergraphConv(32, 3, use_attention=True) def forward(self, x, edge_index, edge_attr): x = self.conv1(x, edge_index, hyperedge_attr=edge_attr) x = F.relu(x) x = F.dropout(x, p=0.5, training=self.training) return self.conv2(x, edge_index, hyperedge_attr=edge_attr)训练就是标准循环,一步反向传播写出来如下:
model = HyperGNN() opt = torch.optim.Adam(model.parameters(), lr=0.01) loss_fn = torch.nn.CrossEntropyLoss() model.train() opt.zero_grad() loss = loss_fn(model(data.x, data.edge_index, edge_attr), data.y) loss.backward() opt.step()如果任务只需要朴素超图卷积,把use_attention去掉、不传edge_attr即可,参数更少、也更省显存。
🔍 适用边界:哪些任务不必硬上超图
上手前建议先核对三点。
- 高阶语义是否真实存在。如果"群"只是若干两两关系的集合,普通图加边特征就能表达,超边维度反而增加参数和内存开销。
- 超边规模。聚合成本由所有超边尺寸之和(即 edge_index 的元素数)决定,一条超边里塞上几千个节点时,scatter 开销线性上涨,需要先做采样或粗粒度压缩。
- 如果任务主体仍是节点两两之间的交互,PyG 更通用的"图结构 + Transformer"路线值得优先考虑,思路是把图结构编码成空间/边信息再送入注意力层:
下一步可以做什么
- 阅读
torch_geometric/data/hypergraph_data.py中subgraph的实现,理解采样节点时超边如何保留、如何重标号,这是大规模超图训练的基础。 - 跑一遍
test/nn/conv/test_hypergraph_conv.py,其中"节点多于超边"和"超边多于节点"两个用例,很适合用来验证自己数据在形状上的正确性。 - 如果数据带时间维度,参考
torch_geometric/datasets/cornell.py(CornellTemporalHyperGraphDataset)里超边随时间演化的数据组织方式。
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考