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

是否存在可提取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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.27 17:45:01