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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.31 22:48:22