PyTorch中torch.unique(带return_count参数)的逆操作实现方法
PyTorch中torch.unique(return_count=True)的逆操作实现
存在对应的逆操作,最直接的实现方式是使用PyTorch内置的torch.repeat_interleaveAPI。
核心用法
torch.repeat_interleave的作用是将输入张量的每个元素重复指定次数,正好匹配torch.unique返回的「唯一值序列+对应出现次数」的还原需求,示例如下:
- 输入:去重后的唯一值张量
unique_vals = torch.tensor([1,2,3]),各值出现次数张量counts = torch.tensor([2,2,2]) - 调用逆操作:
torch.repeat_interleave(unique_vals, counts) - 输出:
tensor([1, 1, 2, 2, 3, 3]),和torch.unique的原始输入完全一致
完整验证代码
import torch # 原始输入张量 raw_tensor = torch.tensor([1, 1, 2, 2, 3, 3]) # 执行torch.unique并开启return_count unique_values, count_list = torch.unique(raw_tensor, return_counts=True) # 执行逆操作还原 restored_tensor = torch.repeat_interleave(unique_values, count_list) # 验证还原结果和原始输入是否一致 print(torch.equal(raw_tensor, restored_tensor)) # 输出为 True
注意事项
- 该还原方式默认要求你调用
torch.unique时使用默认的sorted=True参数:如果当初调用torch.unique时设置了sorted=False,返回的唯一值序列顺序和原始输入中元素首次出现的顺序一致,此时用repeat_interleave还原的序列顺序也会和原始输入对齐。 - 如果你需要完全还原原始输入的元素顺序不受
sorted参数影响,可以在调用torch.unique时额外开启return_inverse=True,得到原始每个元素对应唯一值列表的索引,直接用unique_values[inverse_indices]即可100%还原原始输入序列,示例如下:
raw_tensor = torch.tensor([3,1,1,2,2,3]) unique_values, inverse_indices, count_list = torch.unique(raw_tensor, return_counts=True, return_inverse=True) # 用inverse_indices还原,完全匹配原始顺序 restored_tensor = unique_values[inverse_indices] print(torch.equal(raw_tensor, restored_tensor)) # 输出为True
内容的提问来源于stack exchange,提问作者Exia
相关产品推荐
相关产品推荐

