如何基于Mask向量化分割1D PyTorch张量为张量列表?
按掩码分割1D PyTorch张量的向量化实现
我有一个1D数据张量A,以及和A形状相同的掩码M,示例掩码如下:
tensor([0,0,0,0,1,1,2,2,2,3,3,3,3,3])
我需要根据掩码M把1D张量A分割成1D张量的列表。用循环可以轻松实现,但效率不高,现在想找向量化的解决方案。以下是循环实现的示例:
>>> A = torch.tensor([1,2,3,4,5,6,7,8,9]) >>> M = torch.tensor([0,0,1,1,1,2,2,2,3]) >>> [A[M==i] for i in torch.unique(M)] >>> [tensor([1, 2]), tensor([3, 4, 5]), tensor([6, 7, 8]), tensor([9])]
高效向量化实现方案
根据需求,以下是几种无需循环的高效实现方式,适配不同场景:
场景1:掩码M是连续分组的(如示例中的[0,0,1,1,1,...])
这种情况可以直接通过检测掩码的变化点计算分割长度,无需排序,效率最高:
import torch A = torch.tensor([1,2,3,4,5,6,7,8,9]) M = torch.tensor([0,0,1,1,1,2,2,2,3]) # 检测掩码中值发生变化的位置(索引+1) change_indices = torch.where(M[1:] != M[:-1])[0] + 1 # 计算每个分组的长度 if len(change_indices) == 0: split_lengths = [len(M)] else: split_lengths = torch.cat([ change_indices[:1], change_indices[1:] - change_indices[:-1], torch.tensor([len(M) - change_indices[-1]]) ]).tolist() # 分割张量 result = torch.split(A, split_lengths) print(result) # 输出:(tensor([1, 2]), tensor([3, 4, 5]), tensor([6, 7, 8]), tensor([9]))
场景2:掩码M是无序的(如[0,1,0,2,1]),需保留原张量元素顺序
这种情况先对张量按掩码值稳定排序,再分割,既保证效率又保留元素原顺序:
import torch A = torch.tensor([1,3,2,5,4]) M = torch.tensor([0,1,0,2,1]) # 按掩码值稳定排序,相同掩码的元素保留原顺序 sorted_indices = torch.argsort(M, stable=True) sorted_A = A[sorted_indices] sorted_M = M[sorted_indices] # 检测排序后掩码的变化点 change_indices = torch.where(sorted_M[1:] != sorted_M[:-1])[0] + 1 # 计算分割长度 split_lengths = torch.cat([ change_indices[:1], change_indices[1:] - change_indices[:-1], torch.tensor([len(sorted_M) - change_indices[-1]]) ]).tolist() # 分割得到结果 result = torch.split(sorted_A, split_lengths) print(result) # 输出:(tensor([1, 2]), tensor([3, 4]), tensor([5]))
通用方案(兼容所有掩码情况)
如果不确定掩码是否连续,可通过torch.unique获取分组计数,结合稳定排序实现:
import torch A = torch.tensor([1,3,2,5,4]) M = torch.tensor([0,1,0,2,1]) # 获取掩码唯一值及对应元素数量 unique_vals, counts = torch.unique(M, return_counts=True) # 稳定排序后分割 sorted_indices = torch.argsort(M, stable=True) sorted_A = A[sorted_indices] result = torch.split(sorted_A, counts.tolist()) print(result) # 输出:(tensor([1, 2]), tensor([3, 4]), tensor([5]))
这些方法都避免了循环遍历每个掩码值,利用PyTorch内置向量化操作提升效率,张量规模越大,性能提升越显著。
内容的提问来源于stack exchange,提问作者sebko_iic
相关产品推荐
相关产品推荐

