PyG推荐系统如何批量将预测矩阵中已有边对应值置零
PyG推荐系统已交互课程置零的无循环实现
你不需要写Python循环做逐位置赋值,直接用PyTorch内置的高级索引功能就能高效完成需求,逻辑和你写的循环完全等价,执行效率高几个量级,尤其适配你现在张量放在CUDA设备上的场景,不会出现Python循环带来的设备切换开销。
核心实现代码
一行代码即可完成所有已存在边对应位置的置零:
recs[edge_index[0], edge_index[1]] = 0
实现原理
- 对于形状为
(N, M)的二维张量,PyTorch支持传入两个长度一致的一维张量分别作为行索引、列索引,会一次性定位所有(行索引[i], 列索引[i])对应的坐标点,批量完成赋值操作。 - 整个计算过程完全在PyTorch底层的C++/CUDA runtime执行,没有Python层循环的逐次调度开销,在边数规模较大时性能提升非常明显。
注意事项
- 执行操作前请确认
edge_index和recs在同一个计算设备上,从你给出的代码看二者均在cuda:0,无需额外做设备迁移即可直接运行。 - 上述操作会直接修改原始
recs张量,如果需要保留模型输出的原始预测分数用于其他分析,请先克隆张量再做置零:
# 保留原始预测值的写法 recs_without_hist = recs.clone() recs_without_hist[edge_index[0], edge_index[1]] = 0
- 操作完成后可以通过随机采样校验结果正确性:
# 随机抽取10条历史交互边,验证对应位置分数是否已置0 sample_ids = torch.randint(0, edge_index.size(1), (10,), device=recs.device) check_score = recs[edge_index[0][sample_ids], edge_index[1][sample_ids]] print(check_score) # 输出全为0即表示操作生效
内容的提问来源于stack exchange,提问作者Lucca Baumgärtner
相关产品推荐
相关产品推荐

