PyTorch如何为含重复行的二维张量生成唯一行索引一维张量
PyTorch 1.6 批内多分类重复行target张量实现
现有尺寸为[data_size, token_id_size]的标签张量D,需要为内容完全相同的行分配相同分类索引,两种需求的实现方式如下:
方案1:保留0~(data_size-1)原始索引范围
重复行统一使用该行第一次出现位置的原始索引,重复位置对应的原索引自然空缺,不需要额外做索引映射。
实现代码:
import torch data_size = D.shape[0] # 生成逐行相等判定矩阵:equal_mat[i,j]为True代表D[i]与D[j]内容完全一致 equal_mat = (D.unsqueeze(1) == D.unsqueeze(0)).all(dim=-1) # 逐行取首次匹配到的索引作为target值 target_origin = torch.arange(data_size).unsqueeze(0) \ .masked_select(equal_mat) \ .reshape(data_size, -1)[:, 0] \ .long()
以题目中D的第0、1、8行完全一致、其余行互异的场景为例,该代码输出的target为[0,0,2,3,4,5,6,7,0,9,10,11,12,13,14,15],符合第一种索引规则。
方案2:压缩索引为0~(去重行数-1)连续范围
所有唯一行按规则分配连续索引,重复行复用对应唯一行的索引,索引无空缺。直接使用PyTorch内置的按行去重接口即可实现,性能远高于手动广播对比。
实现代码:
# dim=0指定按行做去重,return_inverse返回每个原始行对应唯一行的索引 # sorted=False表示按行首次出现的顺序分配连续索引,设为True则按张量值排序后分配 _, target_compressed = torch.unique(D, dim=0, return_inverse=True, sorted=False)
同样以题目中的示例场景为例,该代码输出的target为[0,0,1,2,3,4,5,6,0,7,8,9,10,11,12,13],符合第二种压缩索引规则。
性能提示:方案1的广播逻辑会生成尺寸为
[data_size, data_size]的中间布尔矩阵,当batch size超过1024时显存占用会明显升高,大batch场景优先使用方案2。
内容的提问来源于stack exchange,提问作者clement116
相关产品推荐
相关产品推荐

