如何在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
相关产品推荐
相关产品推荐

