如何在PyTorch中使用numpy数组沿指定单维度索引多维张量
解决方案
你当前写法不生效的核心原因分两种情况:
- 如果你只是想重排维度1的顺序,因为你的示例
a是[0,1,2...,9],和维度1的原有索引顺序完全一致,所以输出和原张量没有差异,看起来像没生效;- 如果你要实现对维度0的第i个样本,单独选取维度1的第
a[i]个元素(这是这类索引需求最常见的场景),b[:,a,:]的写法是错误的,该写法会对每个维度0的样本,都选取维度1的全部a对应索引,输出形状和原张量一致,达不到按样本对应选索引的效果。
场景1:重排维度1的顺序
如果你的需求是调整维度1的元素顺序,只需要修改a的元素顺序即可看到效果,示例:
# 把a改成倒序,输出维度1会倒序排列 a = np.array([9,8,7,6,5,4,3,2,1,0]) result = b[:, a, :] # 输出形状:torch.Size([10, 10, 51])
场景2:按维度0的位置对应选取维度1的索引
如果你的需求是每个维度0的样本分别取对应a[i]位置的维度1元素,最终输出形状为[10, 51],有两种常用实现方式:
方法1:用广播索引(最简洁)
给前两个维度分别传入等长的索引数组,自动一一对应选取:
import torch import numpy as np # 你的原始输入 b = torch.randn(10, 10, 51) a = np.array([0,1,2,3,4,5,6,7,8,9]) # 实现索引 result = b[torch.arange(b.shape[0]), a, :] # 输出形状:torch.Size([10, 51])
方法2:用torch.gather(适合更复杂的多维度索引场景)
# 把numpy数组转成torch张量,调整维度和b匹配 a_torch = torch.from_numpy(a).long().unsqueeze(-1).unsqueeze(-1).expand(b.shape[0], 1, b.shape[2]) result = b.gather(dim=1, index=a_torch).squeeze(1) # 输出形状:torch.Size([10, 51])
内容的提问来源于stack exchange,提问作者lalaland
相关产品推荐
相关产品推荐

