如何在PyTorch中高效获取张量唯一行的首次出现索引
如何在PyTorch中高效获取张量唯一行的首次出现索引
嘿,我懂你现在的困扰——用循环逐个查找每个唯一行的首次出现索引,确实在数据量变大的时候会拖慢效率。其实PyTorch本身就提供了更简洁、更高效的原生方法,完全不用写循环就能搞定这个问题!
先说说你现有代码的问题:循环遍历每个唯一行,每次调用torch.where去匹配索引,这种方式在唯一行数量较多时,会累积大量的张量操作开销,效率不够理想。
下面给你一个更PyTorch风格的高效方案,核心是利用torch.scatter_min来一次性获取所有唯一行的首次出现索引:
import torch import numpy as np # 示例张量 data = torch.rand(100, 5) data[np.random.choice(100, 50, replace=False)] = torch.tensor([1.0, 2.0, 3.0, 4.0, 5.0]) # 获取唯一行及相关信息 u_data, inverse_indices, counts = torch.unique(data, dim=0, return_inverse=True, return_counts=True) # 高效获取首次出现索引 # 1. 生成原张量的行索引序列 indices = torch.arange(data.size(0), device=data.device) # 2. 使用scatter_min获取每个唯一行对应的最小(首次出现)索引 unique_indices, _ = torch.scatter_min(indices.unsqueeze(0), 1, inverse_indices.unsqueeze(0)) # 3. 调整为一维张量,和原代码的unique_indices形状一致 unique_indices = unique_indices.squeeze()
方法解释:
indices是一个和原张量行数相同的序列,每个元素对应原张量的行索引(比如第0行对应0,第1行对应1,以此类推)。torch.scatter_min的作用是:根据inverse_indices提供的映射关系,把indices中的值分配到对应的唯一行索引位置上,同时自动保留每个位置的最小值——而首次出现的索引正是每个唯一行对应的最小原索引,这样就一步到位拿到了所有结果。- 最后用
squeeze()把结果从二维(1×N)调整为一维,和你原来的unique_indices结构完全一致。
这个方法的时间复杂度是线性的(和原张量行数成正比),相比循环实现,在数据量较大时效率提升非常明显,而且完全符合PyTorch的张量操作风格,避免了Python循环的额外开销。
备注:内容来源于stack exchange,提问作者ogledala
相关产品推荐
相关产品推荐

