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

如何在PyTorch中使用numpy数组沿指定单维度索引多维张量

解决方案

你当前写法不生效的核心原因分两种情况:

  1. 如果你只是想重排维度1的顺序,因为你的示例a是[0,1,2...,9],和维度1的原有索引顺序完全一致,所以输出和原张量没有差异,看起来像没生效;
  2. 如果你要实现对维度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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.01 01:06:05