PyTorch使用torch.topk后如何筛选满足阈值的张量对应索引
问题背景
现有形状为m×m的张量,为两个张量计算得到的相似度/内积矩阵,需要筛选出所有取值大于0.5的元素对应的原始索引,numpy实现也可接受。
初始测试代码:
import torch x = torch.randn((9052, 512)) similarities = x @ x.T scores, indices = torch.topk(similarities, x.shape[0]) # 取topk等于全量值,返回排序后结果和对应索引
已尝试的无效方案
方案1
mask = torch.ones(scores.size()[0]) mask = 1 - mask.diag() sim_vec = torch.nonzero((scores >= 0.5)*mask)
运行返回形状为[39672595, 2]的张量,不符合预期。
方案2
(scores > 0.5 ).nonzero(as_tuple=True)[0]
运行返回形状为[51152826]的张量,不符合预期。
预期逻辑
返回结果需要和如下伪代码逻辑完全一致:
result = [] for i, row in enumerate(scores): temp = [] for j, value in enumerate(row): if value > 0.5: temp.append(indices[i][j].item()) result.append(temp)
补充说明:曾尝试取矩阵上三角(Upper Triangle)展示元素间相近关系,但未解决阈值筛选、匹配对应原始索引的核心问题,相关代码如下:
import pandas as pd import numpy as np matrix = pd.DataFrame(scores.numpy().astype(np.float32)) upper_tri = matrix.where(np.triu(np.ones(matrix.shape),k=1).astype(np.bool))
错误原因
- 方案1错误:对排序后的
scores矩阵使用原矩阵位置的对角掩码,但topk已经将每行元素按值降序重排,原矩阵的对角元素(样本自匹配值,为行内最大值)已经被移到每行第0位,原j=i位置不再是自匹配值,掩码位置完全错误;且最终取nonzero得到的是排序后的列位置,不是映射后的原始索引。 - 方案2错误:仅提取了
nonzero返回元组的第一个元素(即满足条件的元素行索引),完全没有提取对应列位置和映射后的原始索引,最终得到的是一维行号列表,长度为所有满足条件的元素总数,不符合结构要求。
可直接运行的实现
基于已有topk结果的实现
完全匹配伪代码逻辑,可按需选择是否排除自匹配、是否保留重复对称对:
# 生成阈值筛选掩码 mask = scores > 0.5 # 如需排除样本自匹配(自相似度为行最大值,排在每行第0位),取消注释下一行 # mask[:, 0] = False # 按行提取满足条件的原始索引 result = [] for i in range(scores.shape[0]): row_valid_indices = indices[i][mask[i]].tolist() result.append(row_valid_indices)
更省内存的直接实现(跳过topk步骤)
不需要提前排序,直接在原始相似度矩阵上操作,支持去重(仅保留上三角结果,避免(i,j)和(j,i)重复存储):
# 生成阈值掩码 mask = similarities > 0.5 # 如需排除对角线自匹配,取消注释下一行 # mask.fill_diagonal_(False) # 如需仅保留上三角结果、去掉对称重复对,取消注释下一行 # mask = torch.triu(mask, diagonal=1) # 提取所有满足条件的行列索引 row_ids, col_ids = torch.nonzero(mask, as_tuple=True) # 按行整理为和伪代码一致的结构 result = [col_ids[row_ids == i].tolist() for i in range(similarities.shape[0])]
numpy版本实现
逻辑和上述PyTorch版本一致:
import numpy as np sim_np = similarities.numpy() mask = sim_np > 0.5 # 排除自匹配 np.fill_diagonal(mask, False) # 仅保留上三角去重 mask = np.triu(mask, k=1) row_ids, col_ids = np.nonzero(mask) result = [col_ids[row_ids == i].tolist() for i in range(sim_np.shape[0])]
内容的提问来源于stack exchange,提问作者Deshwal
相关产品推荐
相关产品推荐

