PyTorch中如何为二维张量生成行唯一索引对应的一维张量
PyTorch实现二维张量按行首次出现顺序映射索引
需求说明
对输入二维张量做如下转换:
- 为每一种互不相同的行分配唯一索引,索引值为该行首次出现时的行序号,取值范围为
0到总行数 - 1 - 重复出现的行直接复用已分配的对应索引
转换示例:
# 无重复行场景 [[1,2],[1,3],[1,4]] -> [0,1,2] # 有重复行场景 [[1,2],[1,2],[1,4]] -> [0,0,2] [[1,2],[1,3],[1,2]] -> [0,1,0]
实现方案
直接调用PyTorch内置的torch.unique算子即可实现,核心是指定按行去重、关闭自动排序、返回逆映射索引三个参数,不需要自己写循环遍历,原生算子支持GPU加速,运行效率高。
import torch def get_row_first_idx(input_2d: torch.Tensor) -> torch.Tensor: # 校验输入维度 if input_2d.dim() != 2: raise ValueError("输入必须是二维张量") # dim=0指定按行计算去重 # sorted=False保留元素首次出现的顺序,不会自动排序打乱索引 # return_inverse=True返回原始每行对应去重后类别的索引 _, res = torch.unique( input_2d, dim=0, sorted=False, return_inverse=True ) return res
效果验证
运行测试用例即可验证结果符合要求:
# 测试1 无重复行 t1 = torch.tensor([[1,2],[1,3],[1,4]]) print(get_row_first_idx(t1)) # 输出 tensor([0, 1, 2]) # 测试2 连续重复行 t2 = torch.tensor([[1,2],[1,2],[1,4]]) print(get_row_first_idx(t2)) # 输出 tensor([0, 0, 2]) # 测试3 间隔重复行 t3 = torch.tensor([[1,2],[1,3],[1,2]]) print(get_row_first_idx(t3)) # 输出 tensor([0, 1, 0])
注意事项
- 不要漏写
sorted=False参数,默认情况下torch.unique会对去重结果排序,返回的索引会和首次出现的顺序不一致 - 该实现支持自动微分、GPU张量运算,不需要额外做设备转换,适配所有PyTorch常规工作流
内容的提问来源于stack exchange,提问作者clement116
相关产品推荐
相关产品推荐

