如何基于时间阈值划分PyTorch Geometric的正负训练/测试边?
基于时间阈值划分边并生成正负边索引(替代随机划分工具)
核心思路
先通过时间掩码拆分训练/测试正边,再利用PyG内置的negative_sampling生成对应负边,最后整理成和train_test_split_edges一致的输出格式,无需手动过滤所有Data组件。
实现步骤与代码示例
假设你已经有data(torch_geometric.data.Data实例)和时间阈值time_threshold,且已生成train_mask = (data.edge_time < time_threshold)、test_mask = (data.edge_time >= time_threshold)。
1. 拆分训练/测试正边
直接通过掩码提取对应边的索引和属性:
import torch from torch_geometric.utils import negative_sampling # 拆分正边索引 train_pos_edge_index = data.edge_index[:, train_mask] test_pos_edge_index = data.edge_index[:, test_mask] # 若有边时间/属性,同步拆分(按需保留) train_edge_time = data.edge_time[train_mask] test_edge_time = data.edge_time[test_mask]
2. 生成训练/测试负边
利用negative_sampling生成符合要求的负边,注意避免命中已有正边:
- 训练负边:仅需排除训练集内的正边(保证训练阶段只用到训练时间窗口内的边信息)
- 测试负边:需排除所有正边(训练+测试),避免泄露测试集正边信息
# 生成训练负边 train_neg_edge_index = negative_sampling( edge_index=train_pos_edge_index, num_nodes=data.num_nodes, num_neg_samples=train_pos_edge_index.size(1), # 与训练正边数量一致 directed=False # 若为有向图,设为True ) # 生成测试负边:先合并所有正边,避免负边命中任何正样本 all_pos_edge_index = torch.cat([train_pos_edge_index, test_pos_edge_index], dim=1) test_neg_edge_index = negative_sampling( edge_index=all_pos_edge_index, num_nodes=data.num_nodes, num_neg_samples=test_pos_edge_index.size(1), directed=False )
3. 整理成兼容原API的Data结构
将拆分得到的正负边整合到Data对象中,和train_test_split_edges的输出格式对齐,方便后续模型调用:
# 克隆原数据,避免修改原始数据 split_data = data.clone() # 替换训练用的边索引(若模型需要用训练边作为输入) split_data.edge_index = train_pos_edge_index split_data.edge_time = train_edge_time # 添加正负边索引,与原API输出字段一致 split_data.train_pos_edge_index = train_pos_edge_index split_data.train_neg_edge_index = train_neg_edge_index split_data.test_pos_edge_index = test_pos_edge_index split_data.test_neg_edge_index = test_neg_edge_index # 若需要验证集,可新增val_mask拆分后重复上述步骤,添加val_*字段
关键注意事项
- 如果你的场景是动态图(节点随时间新增),生成负边时需限制节点范围:比如训练负边只能用
time_threshold前已出现的节点,可通过节点的时间属性生成train_node_mask,再传入negative_sampling的subset参数。 - 若需要自定义负边生成逻辑(如基于时间的约束),可以在
negative_sampling的基础上二次过滤,比如筛选负边的两个节点的时间都早于对应阈值。
内容的提问来源于stack exchange,提问作者zelda26
相关产品推荐
相关产品推荐

