PyTorch中如何匹配张量索引与列表值,实现正确实体ID映射?
解决PyTorch链接预测中实体ID与名称的正确映射问题
问题根源
用nonzero()筛选张量后,直接用筛选结果的新数组索引去查inv_entity_dict会出错——因为inv_entity_dict映射的是原张量的实体ID(即原张量的位置索引),而非筛选后数组的下标。比如你提到的Sirolimus,它的原ID不在nonzero()返回的索引里,但误用筛选后数组的下标去查,就会错误匹配到对应ID的实体。
正确实现步骤
假设你有以下变量:
inv_entity_dict: 键为实体ID(整数)、值为实体名称的字典raw_tensor: 原始预测概率张量nonzero_indices:raw_tensor.nonzero()得到的原张量中非零元素的索引
提取原实体ID
把nonzero()返回的索引转换成原实体ID的列表:# 若raw_tensor是一维张量,nonzero返回[N,1]形状的张量,需压缩为一维 original_entity_ids = nonzero_indices.squeeze().tolist()映射为实体名称
直接用提取到的原实体ID查询inv_entity_dict,而非筛选后概率数组的下标:# 筛选得到对应概率值 filtered_probs = raw_tensor[nonzero_indices].squeeze().tolist() # 匹配对应实体名称 entity_names = [inv_entity_dict[entity_id] for entity_id in original_entity_ids]关联概率与实体名称
如需将概率和实体名称对应,可打包成元组列表:prob_entity_pairs = list(zip(filtered_probs, entity_names))
错误示例对比
错误做法(误用筛选后数组的索引):
# 错误:i是筛选后数组的下标,并非原实体ID filtered_probs = raw_tensor[nonzero_indices].squeeze().tolist() wrong_entity_names = [inv_entity_dict[i] for i in range(len(filtered_probs))]
这种写法会把筛选后数组的第0个元素对应到ID=0的实体,完全忽略原张量中的真实ID,必然导致匹配错误。
内容的提问来源于stack exchange,提问作者Ssong
相关产品推荐
相关产品推荐

