如何基于另一数组的唯一索引拆分PyTorch多维张量?
基于PyTorch张量的维度拆分问题解决
问题描述
现有两个PyTorch张量a和b:
import torch torch.manual_seed(0) # 保证可复现 a = torch.rand(size = (5, 10, 1)) b = torch.tensor([3, 3, 1, 5, 3, 1, 0, 2, 1, 2])
想要基于b里的唯一值,拆分a的第2维度(也就是dim=1),预期得到一个列表:列表元素个数等于b中唯一值的数量,每个元素是形状为(5, 对应唯一值的元素个数, 1)的张量。
试过的代码:
# 获取b的唯一值和唯一索引 unique_values, unique_indices = torch.unique(b, return_inverse = True) # 基于唯一索引拆分a的dim=1维度 l = torch.tensor_split(a, unique_indices, dim = 1)
结果里出现了tensor([], size=(5, 0, 1))这类空张量,需要搞清楚原因,以及怎么实现正确的需求。
问题原因
torch.unique(..., return_inverse=True)返回的unique_indices是原张量b中每个元素对应的唯一值索引,比如这个例子里它是tensor([3, 3, 1, 4, 3, 1, 0, 2, 1, 2])。但torch.tensor_split的第二个参数需要的是切分的位置列表——也就是你要从哪些维度索引处把张量切开。直接把unique_indices传进去,相当于让函数在这些索引位置挨个切分,很多位置连续重复或者顺序混乱,自然就切出空张量了。
正确实现方法
方法1:遍历唯一值,筛选对应索引切片
最直观的方式就是遍历b的每个唯一值,找出该值在b里的所有位置,再用这些位置对a的dim=1维度做切片:
import torch torch.manual_seed(0) a = torch.rand(size=(5, 10, 1)) b = torch.tensor([3, 3, 1, 5, 3, 1, 0, 2, 1, 2]) unique_values = torch.unique(b) result = [] for val in unique_values: # 找出b中等于当前唯一值的所有索引 idx = torch.where(b == val)[0] # 按索引切片a的dim=1维度 sliced_tensor = a[:, idx, :] result.append(sliced_tensor) # 验证结果 for i, tensor in enumerate(result): print(f"唯一值{unique_values[i]}对应的张量形状:{tensor.shape}")
输出:
唯一值0对应的张量形状:torch.Size([5, 1, 1]) 唯一值1对应的张量形状:torch.Size([5, 3, 1]) 唯一值2对应的张量形状:torch.Size([5, 2, 1]) 唯一值3对应的张量形状:torch.Size([5, 3, 1]) 唯一值5对应的张量形状:torch.Size([5, 1, 1])
方法2:分组排序后拆分(更高效)
如果要处理大规模张量,这个方法更高效:先按unique_indices给a排序,再找出分组的边界位置做拆分:
import torch torch.manual_seed(0) a = torch.rand(size=(5, 10, 1)) b = torch.tensor([3, 3, 1, 5, 3, 1, 0, 2, 1, 2]) unique_values, unique_indices = torch.unique(b, return_inverse=True) # 按unique_indices排序,得到排序后的索引 sorted_idx = torch.argsort(unique_indices) # 对a做排序,让同一组的元素集中在一起 sorted_a = a[:, sorted_idx, :] # 找出分组的边界(每个分组结束的下一个位置) sorted_inverse = unique_indices[sorted_idx] split_pos = torch.where(sorted_inverse[1:] != sorted_inverse[:-1])[0] + 1 # 拆分得到结果 result = torch.tensor_split(sorted_a, split_pos, dim=1) # 验证结果 for i, tensor in enumerate(result): print(f"唯一值{unique_values[i]}对应的张量形状:{tensor.shape}")
这个方法的结果和方法1完全一致,但避免了循环里的多次切片,处理大张量时速度更快。
内容的提问来源于stack exchange,提问作者bird
相关产品推荐
相关产品推荐

