You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.31 00:39:17