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

如何通过一次操作从PyTorch多维张量中按指定索引提取元素?

解决方案

可以通过高级索引直接对原张量执行一次操作得到目标结果,核心是构造匹配需求的索引组合:

方式一:直接拼接成一维序列(对应分步操作后的拼接结果)

import torch

arr = torch.randint(0, 9, (100, 50, 3))
# 构造第一维度索引:重复每个目标样本的索引,次数等于要提取的元素数量
idx_dim0 = torch.tensor([5]*6 + [55]*6)
# 构造第二维度索引:分别对应两个样本的提取范围(左闭右开,5-10对应5:11,10-15对应10:16)
idx_dim1 = torch.tensor([5,6,7,8,9,10] + [10,11,12,13,14,15])
# 一次索引得到最终结果
final_result = arr[idx_dim0, idx_dim1]

最终结果形状为 torch.Size([12, 3]),和分步操作(先取partial_arr、分别切片后拼接)的输出完全一致。

方式二:保留(2,6,3)的样本维度结构

如果需要保留两个样本的独立维度,可以通过掩码索引实现:

import torch

arr = torch.randint(0, 9, (100, 50, 3))
indices = torch.tensor([5, 55])
# 定义每个样本的切片起止位置
starts = torch.tensor([5, 10])
ends = torch.tensor([11, 16])

# 生成第二维度的全范围索引,并匹配每个样本的提取范围
idx_dim1 = torch.arange(50).unsqueeze(0).repeat(2, 1)
mask = (idx_dim1 >= starts.unsqueeze(1)) & (idx_dim1 < ends.unsqueeze(1))

# 一次索引并整理形状
final_result = arr[indices.unsqueeze(1).repeat(1,6), idx_dim1[mask]].reshape(2,6,3)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 18:12:33