PyTorch:从张量第二维度无放回选指定索引及高效实现与维度疑问
问题解答
一、更高效的实现方式
你的现有实现逻辑是对的,但针对均匀无放回采样的场景,用torch.randperm替代torch.multinomial能获得更高效率——不需要计算多项式分布的采样过程,直接生成随机排列后截取前N个元素即可:
# 生成128个元素的随机排列,重复64次适配第一维度 perm = torch.randperm(128, device=emb_user.device).repeat(64, 1) # 截取每个样本对应的前16个索引 idx = perm[:, :16] # 索引采样后的嵌入 sampled_emb_user = emb_user[torch.arange(64).unsqueeze(-1), idx]
如果你的采样权重不是均匀分布,torch.multinomial是必要选择;但均匀采样场景下,randperm的速度会更快,尤其当第二维度长度较大时。
另外,原代码的索引部分也可以用torch.gather改写,逻辑等价但可读性稍好:
idx = torch.multinomial(torch.ones(64, 128, device=emb_user.device), 16) sampled_emb_user = emb_user.gather(1, idx.unsqueeze(-1).expand(-1, -1, 256))
二、为什么torch.multinomial不支持超过2维的权重矩阵
PyTorch官方对torch.multinomial的设计仅接受1D或2D的权重张量:
- 1D张量:对应单个概率分布,执行n次采样;
- 2D张量:每一行对应一个独立的概率分布,对每一行分别执行n次采样。
如果传入更高维度的张量(比如3D),函数无法明确判断哪一维度对应「概率分布维度」——是最后一维?还是中间某一维?官方未做更高维的适配,因此直接报错。
若需处理更高维的权重,可先将张量reshape为2D,采样完成后再reshape回原维度。比如权重是(64, 10, 128),可这样处理:
weights = torch.ones(64, 10, 128) # 转换为2D张量:(64*10, 128) weights_2d = weights.reshape(-1, 128) idx = torch.multinomial(weights_2d, 16) # 将采样索引还原为原维度:(64, 10, 16) idx = idx.reshape(64, 10, 16)
内容的提问来源于stack exchange,提问作者sachinruk
相关产品推荐
相关产品推荐

