如何在tensor中检索指定值、整行/整列/部分匹配数据的索引
PyTorch张量匹配检索索引方案
具体问题解法:查找[7,5]对应的行索引
首先定义示例张量,通过逐元素匹配+按行全匹配判断拿到掩码,最后提取索引:
import torch t = torch.tensor([[6, 6], [4, 8], [7, 5], [7, 4], [6, 4]]) target = torch.tensor([7,5]) # 生成行匹配掩码:行内所有元素和目标相等则为True row_mask = (t == target).all(dim=1) # 提取匹配的行索引 row_index = row_mask.nonzero(as_tuple=True)[0] print(row_index) # 输出 tensor([2]),对应行索引为2
如果需要将返回的张量索引转为Python原生列表,调用.tolist()方法即可。
通用场景检索方案
1. 整行匹配
逻辑和上述具体问题一致:指定匹配维度为行维度(dim=1),判断该行所有元素都和目标行相等,多个匹配行也会全部返回,适合批量检索场景。
2. 整列匹配
和整行匹配逻辑对称:指定匹配维度为列维度(dim=0),判断该列所有元素都和目标列相等
示例:查找值为[6,8,5,4,4]的列索引
target_col = torch.tensor([6,8,5,4,4]) # 对目标列升维后和原张量做广播匹配,按列判断全匹配 col_mask = (t == target_col.unsqueeze(1)).all(dim=0) col_index = col_mask.nonzero(as_tuple=True)[0] print(col_index) # 输出 tensor([1]),对应第二列(索引从0开始)
3. 行列部分匹配
不需要整行/整列全匹配,仅判断指定位置的元素满足条件即可,适配灵活的检索需求:
示例1:查找第一列值为7的所有行索引
# 仅判断第一列(列索引为0)等于7 partial_mask = t[:, 0] == 7 partial_index = partial_mask.nonzero(as_tuple=True)[0] print(partial_index) # 输出 tensor([2, 3]),对应行索引2和3
示例2:查找第二列值大于等于5的所有行索引
partial_mask = t[:, 1] >=5 partial_index = partial_mask.nonzero(as_tuple=True)[0] print(partial_index) # 输出 tensor([0, 1, 2])
内容的提问来源于stack exchange,提问作者sten
相关产品推荐
相关产品推荐

