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

如何在PyTorch张量的axis-1轴上按索引逐样本选取元素?

解决PyTorch中按样本对应索引选取张量行的问题

给定形状为(2,3,4)的PyTorch张量:

import torch
x = torch.tensor([[[-0.9118,  1.4676, -0.4684, -0.6343],
         [ 1.5649,  1.0218, -1.3703,  1.8961],
         [ 0.8652,  0.2491, -0.2556,  0.1311]],
        [[ 0.5289, -1.2723,  2.3865,  0.0222],
         [-1.5528, -0.4638, -0.6954,  0.1661],
         [-1.8151, -0.4634,  1.6490,  0.6957]]])

需要沿axis=1维度,为每个样本选取对应行:比如索引张量indices = torch.tensor([0, 2]),要求从x[0]取第0行,x[1]取第2行,得到形状为(2,1,4)的输出。

torch.index_select(x, 1, indices)之所以不符合需求,是因为它会对每个样本都应用整个索引列表,最终得到形状(2,2,4)的结果(每个样本都包含第0和第2行),而非我们需要的逐样本单一行选取。

方法1:高级索引

通过构造样本维度的索引,结合给定的行索引实现逐样本选取:

indices = torch.tensor([0, 2])
# 构造样本索引:[0, 1]
sample_indices = torch.arange(x.shape[0])
# 选取对应元素后,通过unsqueeze增加维度到(2,1,4)
result = x[sample_indices, indices].unsqueeze(1)

输出结果:

tensor([[[-0.9118,  1.4676, -0.4684, -0.6343]],
        [[-1.8151, -0.4634,  1.6490,  0.6957]]])

方法2:torch.gather

torch.gather专门用于按指定维度和索引收集元素,需要将索引张量调整为与输入张量匹配的形状:

indices = torch.tensor([0, 2])
# 将indices调整为(2,1,1),再扩展到(2,1,4)以匹配x的最后一维
expanded_indices = indices.unsqueeze(1).unsqueeze(2).expand(-1, -1, x.shape[2])
result = torch.gather(x, dim=1, index=expanded_indices)

输出结果和方法1一致,形状为(2,1,4)。

内容的提问来源于stack exchange,提问作者psuresh

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 11:15:44