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

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)

代码解释

  1. 标记变化位置:通过x[:,1:] != x[:,:-1]找出每行中当前元素与前一个不同的位置,开头补True确保第一个元素被识别为新组。
  2. 生成组ID:用cumsum对变化标记累加,得到每个元素所属的组ID,同组元素ID连续且唯一。
  3. 计算组内计数:用全局索引减去组起始索引,自然得到组内从0开始的计数序列。

这种方法完全基于PyTorch矢量化操作,避免了Python循环,处理大张量时效率更高。

内容的提问来源于stack exchange,提问作者D V

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.18 07:58:11