PyTorch中如何逐行统计张量内唯一元素的出现次数
PyTorch排序张量逐行元素组内计数实现
输入已排序PyTorch张量:
import torch input_tensor = torch.tensor([ [0, 0, 0, 1, 1, 1, 1, 3, 3, 5, 5, 5, 5], [2, 2, 2, 2, 3, 3, 4, 5, 5, 5, 6, 6, 6] ])
需要得到的结果:
target_tensor = torch.tensor([ [0, 1, 2, 0, 1, 2, 3, 0, 1, 0, 1, 2, 3], [0, 1, 2, 3, 0, 1, 0, 0, 1, 2, 0, 1, 2] ])
核心需求
对每行中连续重复的元素,从0开始逐个计数,每组相同元素的计数序列为0,1,2,...,n-1(n为该组元素数量)。
高效矢量化实现
无需逐行循环,用PyTorch内置操作完成:
import torch # 输入张量 x = torch.tensor([ [0, 0, 0, 1, 1, 1, 1, 3, 3, 5, 5, 5, 5], [2, 2, 2, 2, 3, 3, 4, 5, 5, 5, 6, 6, 6] ]) # 1. 标记每行中元素发生变化的位置,开头补True(第一个元素视为新组起点) diff = torch.cat([torch.ones((x.shape[0], 1), dtype=torch.bool), x[:, 1:] != x[:, :-1]], dim=1) # 2. 生成每行元素的组ID,同组元素ID相同 group_ids = torch.cumsum(diff.int(), dim=1) - 1 # 3. 生成每行的索引序列 arange = torch.arange(x.shape[1], device=x.device).unsqueeze(0).repeat(x.shape[0], 1) # 4. 计算每个组的起始索引 start_indices = torch.zeros_like(x) start_indices.scatter_(1, torch.where(diff)[1].view(x.shape[0], -1), arange[diff].view(x.shape[0], -1)) group_start = torch.cummax(start_indices, dim=1)[0] # 5. 每个位置的计数 = 当前索引 - 组起始索引 result = arange - group_start print(result)
代码解释
- 标记变化位置:通过
x[:,1:] != x[:,:-1]找出每行中当前元素与前一个不同的位置,开头补True确保第一个元素被识别为新组。 - 生成组ID:用
cumsum对变化标记累加,得到每个元素所属的组ID,同组元素ID连续且唯一。 - 计算组内计数:用全局索引减去组起始索引,自然得到组内从0开始的计数序列。
这种方法完全基于PyTorch矢量化操作,避免了Python循环,处理大张量时效率更高。
内容的提问来源于stack exchange,提问作者D V
相关产品推荐
相关产品推荐

