尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

PyG超图实战指南:HyperGraphData 与 HypergraphConv 从零到跑通

PyG超图实战指南:HyperGraphData 与 HypergraphConv 从零到跑通 PyG超图实战指南HyperGraphData 与 HypergraphConv 从零到跑通【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric推荐系统里用户-商品-标签这种三元关系用普通图建模时只能拆成若干条二元边而拆分方式本身就引入了信息失真用户和标签之间的关联是靠商品中转出来的并不存在一条真实的直接联系。这类天然以一组节点同时成立为语义的关系适合用超图hypergraph来承载——超边hyperedge是一条能同时连接任意多个节点的边一个超边就对应一个完整的高阶关系。PyTorch GeometricPyG提供了对应的数据结构HyperGraphData和卷积层HypergraphConv两者配合可以直接完成超图卷积网络Hypergraph Convolutional Network, HGCN的建模与训练。 一个三元关系为什么二元边不够用结论先说当关系单元是一个组而不是一对时先建模组再谈组内关系比硬拆成对更忠实于数据。拿上面的三元关系举例一条超边{用户u, 商品c, 标签t}表示u 因为 t 而购买 c这三个节点必须同时在场这条关系才成立。若拆成(u,c)、(c,t)、(u,t)三条二元边(u,t)这条是虚构出来的——用户从没直接买过标签。这种组即语义的模式在多个领域反复出现分子里一个官能团羟基、氨基天然由多个原子共同构成做分子性质预测时把它当作一条超边比两两连键更贴近化学事实群聊、多人共同签署的文档同理。区别只在于二元关系图建模的是谁和谁有边超图建模的是哪些节点构成一个整体。 最小可运行示例从构造到训练跑通PyG 中用HyperGraphData承载超图数据源码见 torch_geometric/data/hypergraph_data.py用HypergraphConv做特征聚合源码见 torch_geometric/nn/conv/hypergraph_conv.py。下面这段代码包含数据构造、两层模型与完整训练循环可以直接运行import torch import torch.nn.functional as F from torch_geometric.data import HyperGraphData from torch_geometric.nn import HypergraphConv # 数据6个节点超边0连接{0,1,2}超边1连接{1,2,3,4} x torch.randn(6, 16) y torch.tensor([0, 0, 1, 1, 2, 2]) edge_index torch.tensor([ [0, 1, 2, 1, 2, 3, 4], # 第一行节点索引 [0, 0, 0, 1, 1, 1, 1], # 第二行超边编号 ]) data HyperGraphData(xx, edge_indexedge_index, yy) print(data.num_nodes, data.num_edges) # 6 2 class HGNN(torch.nn.Module): def __init__(self): super().__init__() self.conv1 HypergraphConv(16, 32, use_attentionTrue) self.conv2 HypergraphConv(32, 3, use_attentionTrue) def forward(self, x, edge_index): x F.relu(self.conv1(x, edge_index)) x F.dropout(x, p0.5, trainingself.training) return F.log_softmax(self.conv2(x, edge_index), dim-1) model HGNN() opt torch.optim.Adam(model.parameters(), lr0.01) loss_fn torch.nn.CrossEntropyLoss() for epoch in range(1, 101): model.train() opt.zero_grad() loss loss_fn(model(data.x, data.edge_index), data.y) loss.backward() opt.step() if epoch % 20 0: print(fEpoch {epoch:03d}, Loss: {loss:.4f})节点特征维度是 16、类别数是 3 时两层卷积后接分类即可。若任务不是节点分类把最后的 softmax 换成回归头或池化到图级输出就行前两层结构不用动。 原理拆解超边索引怎么写、卷积公式怎么算理解 PyG 超图实现只需要抓住一件事超边索引用坐标对表示关联矩阵。edge_index形状为[2, C]C是节点-超边关联对的总数不是超边条数第一行列出每个关联对中的节点编号第二行列出它所属的超边编号。上面示例里edge_index[0][3] 1、edge_index[1][3] 1合起来表示节点 1 属于超边 1。编码项普通图Data超图HyperGraphDataedge_index第一行源节点节点编号edge_index第二行目标节点超边编号边数含义edge_index.size(1)edge_index[1].max() 1HyperGraphData的num_edges属性就是按超边编号最大值加 1计算的num_nodes则取节点编号最大值加 1。超边自身的特征放在edge_attr形状[M, D]标量权重放在hyperedge_weight形状[M]M是超边数。卷积层HypergraphConv实现的是 HGCN 论文的算子核心公式只有一个$$\mathbf{X}^{\prime} \mathbf{D}^{-1} \mathbf{H} \mathbf{W} \mathbf{B}^{-1} \mathbf{H}^{\top} \mathbf{X} \mathbf{\Theta}$$各符号一句话解释$\mathbf{H}$关联矩阵指示哪个节点属于哪条超边也就是超边索引的矩阵形式$\mathbf{W}$超边权重标量hyperedge_weight对应的对角矩阵$\mathbf{D}$、$\mathbf{B}$节点与超边的度矩阵起归一化作用避免大超边把特征值刷大$\mathbf{\Theta}$可学习的线性变换参数。实现上的行为与公式一一对应先做线性变换再用 $\mathbf{B}^{-1}\mathbf{H}^{\top}$ 把节点特征上卷到超边节点到超边方向用 $\mathbf{W}\mathbf{H}$ 与 $\mathbf{D}^{-1}$ 把超边信息下卷回节点两个方向各一次propagate。所以一个节点会吸收它所属所有超边的信息一条超边会汇总它包含的所有节点——这正是先建模组、再谈组内关系在计算上的体现。前面提到分子官能团、推荐三元组这类场景落到这个公式里只是 $H$ 中关联对的不同排布方式模型结构完全通用。⚙️ 关键参数与 node/edge 注意力模式怎么选HypergraphConv的默认配置use_attentionFalse, attention_modenode, heads1, concatTrue就是一个不带注意力的 HGCN 层。开注意力时的几个关键参数参数作用选择建议use_attention是否计算注意力系数需要hyperedge_attr参与计算小数据集上默认关闭更稳attention_mode注意力归一化的方向node或edge默认nodeheads/concat多头数 / 是否拼接concatFalse时多头取平均输出维度不变hyperedge_attr超边特征形状[M, D]开注意力时必传否则前向直接抛断言错误hyperedge_weight超边权重形状[M]想区分超边重要性时传入如不同购买行为的强度attention_mode是选择上最有信息量的一个参数两种模式归一化的方向不同node模式在同一条超边内部的节点之间做 softmax。语义是这条超边里的节点谁更重要。适合组内竞争明显的场景比如一个兴趣组内哪些用户/商品更能代表该组。edge模式在同一个节点所属的多条超边之间做 softmax。语义是这个节点身上的多条关系哪条更值得吸收。适合节点普遍挂在多条超边上、且超边间重要性差异大的场景比如一个用户同时属于几十个兴趣组。经验上先跑默认的node模式只在节点跨超边比较成为瓶颈时切到edge。另外注意headsk且concatTrue时输出维度是k * out_channels改完 heads 记得同步下游第一层concatFalse则输出恒为out_channels。⚠️ 适用边界与常见坑先划适用边界超图的价值在于数据里天然存在三元及以上的关系单元。如果你的图其实只是成对关系用普通Data 常规卷积层生态更成熟没必要引入超图反过来硬把超边拆成二元边再建模才是真正会亏的做法。常见坑按踩中概率排序开注意力却不传hyperedge_attr。use_attentionTrue时前向会断言要求超边特征报错信息并不直观构造数据时就顺手生成一个edge_attr能省很多调试时间。两行索引写反。节点行和超边行互换后不会立刻报错但聚合语义完全改变num_nodes/num_edges的推断也会错乱。写完后打印一次两个num_*属性是最便宜的验证手段。把hyperedge_weight和hyperedge_attr搞混。前者是形状[M]的标量权重对应公式里的 $\mathbf{W}$后者是形状[M, D]的特征矩阵只在开注意力时使用——不开注意力时它不参与任何计算传了也是白传。超边规模悬殊。一条超边挂 3 个节点、另一条挂 300 个节点时度归一化虽然会做缩放但大超边汇聚的特征仍是平均感更强的。若业务上超边规模差异巨大可以先按规模对hyperedge_weight做一下标定。混用异构图接口。HyperGraphData的to_heterogeneous()、is_directed()等方法直接抛NotImplementedError超图有自己独立的一套接口别把它当异构图用。 延伸方向动态超图把超边索引按时间切片超边集合随时间演化用于建模会增减成员的兴趣组、随版本变化的文档共现等时序高阶关系。超图 Transformer用 Transformer 的消息传递替换卷积层中的propagate缓解长距离高阶依赖的衰减代价是计算量上升。大规模数据百万节点级超图需要分布式采样与训练可以参考 PyG 分布式训练相关模块与examples/distributed/下的示例组织方式。下一步建议先把上面的最小示例在你自己的三元关系数据上跑一遍确认超边构建合理再逐步加注意力层——数据侧的超边质量比结构侧的调参影响更大。【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表