是否存在可提取N*h*w*2形状PyTorch张量中唯一元组的函数?
PyTorch 提取四维张量最后一维唯一元组及对应原索引实现方法
你之前使用torch.unique()未达到预期效果,是因为该方法默认会将整个张量展平为一维后查找标量唯一值,没有指定按元组维度做去重。具体实现逻辑如下:
- 先把形状为
(N, h, w, 2)的输入张量展平前三个维度,得到形状为(M, 2)的张量(其中M = N * h * w),所有待去重的二元组都会排列在第一维度 - 调用
torch.unique()时指定dim=0按行去重,同时开启返回索引的参数,即可得到唯一元组、对应原展平位置的索引 - 最后通过
torch.unravel_index将展平的一维索引还原为原张量的多维位置索引
完整代码示例
import torch # 示例输入张量 a = torch.tensor([[[[1,2], [2,3], [3,4]], [[4,5], [1,2], [3,4]]]]) print(f"原张量形状:{a.shape}") # 输出:torch.Size([1, 2, 3, 2]) # 展平前三个维度,得到所有二元组的集合 flatten_tuples = a.reshape(-1, 2) # 按行去重,返回唯一元组、原位置映射、唯一元组首次出现的展平索引 unique_tuples, inverse_map, flatten_unique_idx = torch.unique( flatten_tuples, dim=0, return_inverse=True, sorted=True ) print("去重后的唯一二元组:") print(unique_tuples) # 输出: # tensor([[1, 2], # [2, 3], # [3, 4], # [4, 5]]) # 将展平索引还原为原张量的(N, h, w)多维索引 original_unique_idx = torch.unravel_index(flatten_unique_idx, a.shape[:-1]) print("唯一二元组对应原张量的位置索引(依次为N、h、w维度):") print(original_unique_idx) # 输出: # (tensor([0, 0, 0, 0]), tensor([0, 0, 0, 1]), tensor([0, 1, 2, 0])) # 对应四个唯一元组首次出现的位置分别是(0,0,0)、(0,0,1)、(0,0,2)、(0,1,0) # 若需要查看原张量每个位置对应的唯一元组编号,可将映射表还原形状 original_inverse_map = inverse_map.reshape(a.shape[:-1]) print("原张量每个位置对应的唯一元组编号:") print(original_inverse_map) # 输出: # tensor([[0, 1, 2], # [3, 0, 2]])
内容的提问来源于stack exchange,提问作者Xingyu Chen
相关产品推荐
相关产品推荐

