PyTorch Geometric 异构图实战指南从 HeteroConv 到 to_hetero 的完整示例解析【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric导读异构Heterogeneous图是节点与边类型多样的图结构广泛存在于推荐系统、学术网络、知识图谱等场景。本文以 PyTorch GeometricPyG官方示例目录 examples/hetero 中的 14 个可运行脚本为骨架系统讲解异构图的表示方式HeteroData、两种核心建模范式HeteroConv手动逐边类型建模与to_hetero自动转换同构图模型、专用异构图架构HGT、HAN、链接预测、时序图学习、层次采样、CSV 数据构建以及无监督表示学习MetaPath2Vec、DMGI。读完本文你将能够对照源码独立搭建从数据加载、模型设计到训练评估的完整异构图方案。一、异构图的基础HeteroData 与 metadataPyG 中异构图的载体是torch_geometric.data.HeteroData实现见 torch_geometric/data/hetero_data.py。它与同构Data的核心区别在于所有张量都通过**键key**组织为字典形式节点特征x_dict形如{author: tensor, paper: tensor, ...}边索引edge_index_dict键为三元组(源节点类型, 关系名, 目标节点类型)例如(author, writes, paper)标签与掩码按节点类型划分如data[author].train_mask。data.metadata()返回(node_types, edge_types)二元组是后续模型初始化、HeteroConv构建和to_hetero转换所必需的元数据。在 examples/hetero/load_csv.py 中可以看到HeteroData最原生的构建方式data HeteroData() data[user].num_nodes len(user_mapping) # 用户无特征仅声明节点数 data[movie].x movie_x data[user, rates, movie].edge_index edge_index data[user, rates, movie].edge_label edge_label一个值得注意的细节MovieLens 场景下用户节点没有原始特征示例用torch.eye(...)单位矩阵作为user的 one-hot 特征见 hetero_link_pred.py并在赋值后删除num_nodes属性这是因为x的存在会自动推导节点数。二、HeteroConv逐边类型手动建模2.1 核心机制HeteroConv源码见 torch_geometric/nn/conv/hetero_conv.py是一个通用包装器它接收一个{edge_type: MessagePassing层}字典对每条边类型独立执行对应的消息传递层再把指向同一目标节点类型的多条关系的结果按aggr聚合。aggr支持sum、mean、min、max、cat或None默认sum从源码第 13-26 行的group函数可以看出cat走拼接路径其余走torch.stack后按维度归约None则直接堆叠成新维度——因此它尤其适合对不同边类型使用不同消息传递模块的场景。2.2 实战DBLP 节点分类hetero_conv_dblp.pyhetero_conv_dblp.py 展示了完整流程dataset DBLP(path, transformT.Constant(node_typesconference)) data dataset[0]DBLP 数据集含author、paper、term、conference四种节点类型。conference节点原本没有特征通过T.Constant(node_typesconference)为其填充常量特征使消息传递能正常进行。模型定义是HeteroConv的标准用法class HeteroGNN(torch.nn.Module): def __init__(self, metadata, hidden_channels, out_channels, num_layers): super().__init__() self.convs torch.nn.ModuleList() for _ in range(num_layers): conv HeteroConv({ edge_type: SAGEConv((-1, -1), hidden_channels) for edge_type in metadata[1] }) 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: F.leaky_relu(x) for key, x in x_dict.items()} return self.lin(x_dict[author])要点对所有边类型统一使用SAGEConv通过字典推导式一次构建SAGEConv((-1, -1), hidden_channels)使用-1表示惰性lazy初始化输入维度在第一次前向传播时根据实际特征维度自动确定输出时只需取目标节点类型author的表示进行分类模型训练前先用torch.no_grad()前向一次以完成惰性模块的参数初始化脚本第 48-49 行训练循环直接使用data[author].train_mask与data[author].y计算交叉熵损失并在train/val/test三个掩码上分别计算准确率。该脚本是理解异构图 字典化同构图这一思想的最佳起点。三、to_hetero把同构图模型一键转换为异构图模型3.1 核心机制to_hetero(module, metadata, aggrNone)实现在 torch_geometric/nn/to_hetero_transformer.py是 PyG 最具生产力的 API它接收一个针对同构图编写的MessagePassing模块输入输出均为单个张量基于metadata自动将其转换为接收x_dict/edge_index_dict的异构图模块。转换过程中每个边类型获得独立的权重副本指向同一目标节点类型的不同关系的结果按aggr如sum聚合。它对torch.nn.Sequential、Linear、ReLU等常见层均有内建转换规则。3.2 实战ogb-mag 节点分类to_hetero_mag.pyto_hetero_mag.py 是官方文档同款流程的完整版。数据侧ogb-mag数据集的 4 类节点、4 类关系通过T.ToUndirected(mergeTrue)补全反向边以便消息双向传播模型侧先定义同构的Sequential模型再一次性完成转换model Sequential(x, edge_index, [ (SAGEConv((-1, -1), 64), x, edge_index - x), ReLU(inplaceTrue), (SAGEConv((-1, -1), 64), x, edge_index - x), ReLU(inplaceTrue), (Linear(-1, dataset.num_classes), x - x), ]) model to_hetero(model, data.metadata(), aggrsum).to(device)Sequential的x, edge_index - x描述了每个子模块的输入输出签名to_hetero据此识别哪些参数需要按边类型分权、哪些需要按节点类型分权。采样与训练部分同样有代表性通过NeighborLoader或HGTLoader以(paper, train_mask)作为input_nodes采样子图batch_size1024、num_neighbors[10] * 2每批次中只取种子节点部分out[paper][:batch_size]与batch[paper].y[:batch_size]计算损失脚本预留了--use_hgt_loader命令行开关可在邻居采样与 HGT 层次采样之间切换对比两种采样策略。3.3 实战MovieLens 链接预测hetero_link_pred.pyhetero_link_pred.py 展示了 encoder decoder 范式编码器同构的GNNEncoder两层SAGEConv经to_hetero转换后在用户-电影二部图上做消息传递得到z_dict解码器EdgeDecoder拼接(user, movie)两端的嵌入经两层Linear输出预测评分数据切分先用T.ToUndirected()补反向关系(movie, rev_rates, user)再用T.RandomLinkSplit(num_val0.1, num_test0.1, neg_sampling_ratio0.0, edge_types[(user, rates, movie)], rev_edge_types[(movie, rev_rates, user)])做边级切分得到train_data / val_data / test_data损失设计MovieLens 评分分布极不均衡3、4 分多0、1 分极少脚本提供--use_weighted_loss选项用torch.bincount统计评分频数后构造weight.max() / weight的反频权重配合自定义weighted_mse_loss训练评估指标为 RMSE预测值clamp(min0, max5)。四、专用异构图架构HGT 与 HAN4.1 HGTHeterogeneous Graph Transformerhgt_dblp.pyexamples/hetero/hgt_dblp.py 使用HGTConv实现于 torch_geometric/nn/conv/hgt_conv.py在 DBLP 上做节点分类。HGT 的核心思想是每个元关系(源类型, 关系, 目标类型)学习独立的注意力计算通过相对时间编码与类型相关的权重变换建模不同类型节点间的异构交互。脚本结构清晰先用ModuleDict为每种节点类型配置一个Linear(-1, hidden_channels)做类型专属的输入投影再堆叠num_layers层HGTConv最后仅取author的表示过Linear分类。模型参数为hidden_channels64, out_channels4, num_heads2, num_layers1训练 100 个 epoch。HGT 无需显式构造 metapath而是让注意力机制在元关系层面自动学习适合关系语义丰富的图。4.2 HANHeterogeneous Graph Attention Networkhan_imdb.pyexamples/hetero/han_imdb.py 展示 HANHANConv实现于 torch_geometric/nn/conv/han_conv.py的用法。HAN 是基于 metapath的注意力模型先在每个 metapath 内部做节点级注意力聚合再对不同 metapath 做语义级注意力加权。数据侧使用T.AddMetaPaths变换metapaths [[(movie, actor), (actor, movie)], [(movie, director), (director, movie)]] transform T.AddMetaPaths(metapathsmetapaths, drop_orig_edge_typesTrue, drop_unconnected_node_typesTrue)这里两条 metapath 分别表达电影-演员-电影与电影-导演-电影两种语义关联drop_orig_edge_typesTrue表示丢弃原始边类型仅保留 metapath 生成的新边drop_unconnected_node_typesTrue则剔除在 metapath 中无关联的节点类型。训练循环带早停patience100验证集准确率连续 100 个 epoch 不提升即停止。五、层次采样与内存优化hierarchical_sage.py大规模异构图如 ogb-mag 包含数十万节点无法全图前向。hierarchical_sage.py 演示了NeighborLoader采样 trim_to_layer来自 torch_geometric/utils/trim_to_layer.py的分层训练方案for i, conv in enumerate(self.convs): x_dict, edge_index_dict, _ trim_to_layer( layeri, num_sampled_nodes_per_hopnum_sampled_nodes_dict, num_sampled_edges_per_hopnum_sampled_edges_dict, xx_dict, edge_indexedge_index_dict, ) x_dict conv(x_dict, edge_index_dict) x_dict {key: x.relu() for key, x in x_dict.items()}每层卷积前trim_to_layer根据batch.num_sampled_nodes_dict / num_sampled_edges_dict把超出该层感受野的节点与边裁剪掉避免随层数增加出现采样节点数不降反增的指数膨胀显著降低显存与计算开销。脚本还提供--use-sparse-tensor开关配合T.ToSparseTensor()使用adj_t_dict稀疏张量存储边。该脚本是处理 ogb-mag 等大规模异构数据集的标准范式。六、从原始 CSV 构建异构图load_csv.pyload_csv.py 展示了完全从零开始的流程——不依赖现成数据集类直接下载 MovieLens 的movies.csv与ratings.csv手写三个工具函数load_node_csv(path, index_col, encodersNone)以指定列为索引建立字符串→整数映射对每个指定列应用编码器后拼接成特征矩阵xload_edge_csv(...)把源/目标索引映射为整数生成edge_index可选地对边属性列编码得到edge_attr三种特征编码器SequenceEncoder基于sentence_transformers.SentenceTransformer(all-MiniLM-L6-v2)把电影标题文本编码为稠密嵌入GenresEncoder按|切分类型列做多热multi-hot分类编码IdentityEncoder原样转换为 PyTorch 张量用于评分列。构建完成后与链接预测示例一致地执行两步转换ToUndirected()补反向边并删除反向边上的edge_label避免标签泄漏再用RandomLinkSplit(num_val0.05, num_test0.1, neg_sampling_ratio0.0, ...)切分出train/val/test三份数据。这个示例揭示了HeteroData与 PyG 变换系统组合的底层逻辑适合作为自定义数据集接入的模板。七、无监督表示学习MetaPath2Vec 与 DMGI7.1 MetaPath2Vecmetapath2vec.pymetapath2vec.py 在 AMiner 数据集上训练无监督的MetaPath2Vec实现于 torch_geometric/nn/models/metapath2vec.py。核心是显式定义元路径来引导随机游走metapath [ (author, writes, paper), (paper, published_in, venue), (venue, publishes, paper), (paper, written_by, author), ] model MetaPath2Vec(data.edge_index_dict, embedding_dim128, metapathmetapath, walk_length50, context_size7, walks_per_node5, num_negative_samples5, sparseTrue)参数含义embedding_dim128嵌入维度、walk_length50游走长度、context_size7skip-gram 上下文窗口、walks_per_node5每个节点游走次数、num_negative_samples5负采样数、sparseTrue使用稀疏嵌入。训练使用SparseAdam优化器与model.loader(batch_size128)提供的游走数据加载器。评估方式值得借鉴取model(author, ...)得到作者嵌入仅用 10% 数据训练线性分类器报告 Micro-F1脚本注释标注约 91.8%体现了无监督预训练 下游线性评估的经典范式。7.2 DMGIdmgi_unsup.pydmgi_unsup.py 使用 DMGIDeep Multiplex Graph Infomax实现于 torch_geometric/nn/models/dmgi.py在 IMDB 数据集上学习嵌入。DMGI 面向多关系multiplex异构图对每种关系分别做图卷积编码并计算局部-全局互信息再通过可学习的权重融合各关系视图的嵌入最终以无监督方式产出可用于下游分类的节点表示。八、面向业务的时序链路预测与推荐系统8.1 时序链接预测temporal_link_pred.pyexamples/hetero/temporal_link_pred.py 在 MovieLens 上演示基于时间的边切分与LinkNeighborLoaderperm torch.argsort(data[user, movie].time) train_idx perm[:int(0.8 * perm.size(0))] val_idx perm[int(0.8 * perm.size(0)):int(0.9 * perm.size(0))] test_idx perm[int(0.9 * perm.size(0)):]关键参数time_attrtime指定时间属性、temporal_strategylast表示每个节点只保留最近 K 条时间边、edge_label_timetime[train_idx] - 1把标签边的时间减 1 以确保采样时不会看到未来的监督信息无泄漏。模型同样采用to_hetero转换的SAGEConv编码器加EdgeDecoder解码器以 MSE 损失回归评分、RMSE 评估。8.2 时序推荐系统recommender_system.pyexamples/hetero/recommender_system.py 是该目录下最完整的业务化示例模拟基于历史的用户-商品交互做 Top-K 召回数据预处理仅保留评分 ≥ 4 的高质量交互作为正样本删除评分信息补反向边时序切分按时间取前 80% 为训练边后 20% 为测试边训练LinkNeighborLoader配合neg_samplingdict(modebinary, amount2)二元负采样InnerProductDecoder内积解码binary_cross_entropy_with_logits损失测试分别用NeighborLoader采样得到用户与电影嵌入基于MIPSKNNIndex最大内积 k-NN 索引做高效近邻检索并用EdgeIndex.sparse_narrow排除训练边最后通过 torch_geometric/metrics 提供的LinkPredMAP、LinkPredPrecision、LinkPredRecall计算MAPk、Precisionk、Recallk--k默认 20。8.3 基于 metapath 的二部图模型bipartite_sage.py / bipartite_sage_unsup.py目录中的 bipartite_sage.py 与 bipartite_sage_unsup.py 展示了通过 metapath 训练 GNN 做链接预测的另一条路线利用用户-物品-用户/物品-用户-物品等 metapath 采样邻居进行消息传递。后者将同一方案扩展到大规模 TaoBao 数据集演示了从中小规模实验脚本向工业级数据规模迁移时的工程要点。九、示例全景与选型指南将 examples/hetero/README.md 中的 14 个示例按技术主题归类可得到如下选型地图技术需求推荐示例核心 API / 数据集手动逐边类型建模hetero_conv_dblp.pyHeteroConvSAGEConv/ DBLP同构模型一键转异构to_hetero_mag.pyto_heteroSequential/ ogb-mag异构链接预测评分回归hetero_link_pred.pyto_heteroRandomLinkSplit/ MovieLensTransformer 类异构图模型hgt_dblp.pyHGTConv/ DBLPMetapath 注意力模型han_imdb.pyHANConvAddMetaPaths/ IMDB大规模层次采样训练hierarchical_sage.pyNeighborLoadertrim_to_layer/ ogb-mag从 CSV 自建数据集load_csv.pyHeteroData 自定义编码器 / MovieLens无监督元路径游走metapath2vec.pyMetaPath2VecSparseAdam/ AMiner无监督多关系表示dmgi_unsup.pyDMGI/ IMDB时序链接预测temporal_link_pred.pyLinkNeighborLoadertemporal_strategy/ MovieLens大规模二部图链接预测bipartite_sage_unsup.pymetapath 采样 / TaoBao时序推荐系统Top-K 召回recommender_system.pyMIPSKNNIndex 检索指标 / MovieLens十、最佳实践小结从这些示例中可以提炼出 PyG 异构图的通用工程要点数据准备三件套无特征节点用T.Constant填充或用torch.eye构造 one-hot用T.ToUndirected()补反向边保证消息双向传播链接预测任务务必在补反向边之后用RandomLinkSplit切分并删除反向边上的标签防止泄漏模型选择两条路需要为不同边类型定制不同聚合器时选HeteroConv已有同构图模型时优先to_hetero一行完成转换配合Sequential与惰性初始化-1维度最省事大规模训练NeighborLoader/HGTLoader采样 trim_to_layer剪枝注意input_nodes用(node_type, mask)元组指定种子时序约束带时间属性的数据使用LinkNeighborLoader的time_attr/temporal_strategy与edge_label_time - 1的偏移技巧评估对齐业务回归型链路预测看 RMSE召回型推荐看 MAP/Precision/Recallk且评估时用EdgeIndex.sparse_narrow显式排除训练边。以上所有脚本均可在仓库examples/hetero/目录下直接运行数据自动下载至../../data/是学习与二次开发的最佳起点。【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考