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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.20 11:20:29