如何在PyTorch中获取二维Tensor指定维度的唯一值(无循环)
嘿,这个场景我太熟悉了!循环调用torch.unique处理每行确实在数据量大的时候会拖慢速度,毕竟张量运算的优势就是批量并行处理嘛。给你分享两个无(或极少)循环的实现思路,效率会提升不少:
首先先复现下你的问题场景:
import torch tmp = torch.tensor([[1,2,3,2,4],[0,5,6,7,2],[3,4,5,3,5],[7,5,6,7,7]]) # 预期结果:[[1,2,3,4],[0,5,6,7,2],[3,4,5],[5,6,7]]
方法一:行偏移+全局去重(保留原出现顺序)
这个思路是给每行的元素加上独有的偏移,让不同行的元素不会重复,然后通过全局去重后再还原偏移,最后按行分组。核心操作都是批量张量运算,只有最后分组的小循环(次数等于行数,开销极低):
# 1. 计算每行的偏移量,确保不同行的元素不会重叠 max_val = tmp.max() offsets = torch.arange(tmp.size(0), device=tmp.device) * (max_val + 1) offsets = offsets.unsqueeze(1) # 扩展维度适配广播 # 2. 给每行元素加上偏移,生成全局唯一的元素集合 shifted = tmp + offsets # 3. 全局去重并保留第一次出现的顺序 flattened_shifted = shifted.flatten() unique_shifted, indices = torch.unique(flattened_shifted, return_indices=True) # 按首次出现的索引排序,保证顺序和原张量一致 sorted_indices = indices.sort().values unique_shifted_sorted = flattened_shifted[sorted_indices] # 4. 还原偏移,得到原数值 row_ids = sorted_indices // tmp.size(1) unique_vals = unique_shifted_sorted - offsets.flatten()[row_ids] # 5. 按行分组收集结果(仅此处有循环,次数为行数) result = [] for row_idx in range(tmp.size(0)): result.append(unique_vals[row_ids == row_idx].tolist()) print(result) # 输出:[[1, 2, 3, 4], [0, 5, 6, 7, 2], [3, 4, 5], [5, 6, 7]]
方法二:排序+差异标记(可选择保留原顺序)
如果可以接受唯一值是排序后的,这个方法几乎没有循环,速度最快;如果需要保留原出现顺序,最后只需要少量的行循环处理:
步骤1:获取排序后的唯一值(无循环)
# 对每行排序,然后通过diff标记不同元素 sorted_tmp, _ = tmp.sort(dim=1) # prepend一个不可能出现的数值,确保第一个元素被标记为唯一 diff = torch.diff(sorted_tmp, dim=1, prepend=torch.full((tmp.size(0),1), -1, device=tmp.device)) mask = diff != 0 # 提取每行的唯一值(排序后) sorted_unique = sorted_tmp[mask].split(mask.sum(dim=1).tolist()) result_sorted = [x.tolist() for x in sorted_unique] print(result_sorted) # 输出:[[1, 2, 3, 4], [0, 2, 5, 6, 7], [3, 4, 5], [5, 6, 7]]
步骤2:转换为原出现顺序(少量循环)
如果需要和你预期的原出现顺序一致,可以通过查找每个唯一值在原行中首次出现的索引来排序:
result = [] for i in range(tmp.size(0)): row = tmp[i] unique_vals = sorted_unique[i] # 找到每个唯一值在原行中第一次出现的位置 first_occur_idx = torch.argmax((row == unique_vals.unsqueeze(1)).int(), dim=1) # 按首次出现顺序重新排列唯一值 ordered_unique = unique_vals[first_occur_idx.sort().indices] result.append(ordered_unique.tolist()) print(result) # 输出:[[1, 2, 3, 4], [0, 5, 6, 7, 2], [3, 4, 5], [5, 6, 7]]
为什么这两个方法更快?
原来的循环是对每行单独调用torch.unique,每次都会触发独立的张量运算;而上面的方法把大部分逻辑放在批量张量操作里,充分利用了PyTorch的并行计算优势,行数越多,性能提升越明显。
内容的提问来源于stack exchange,提问作者bob wong
相关产品推荐
相关产品推荐

