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

PyTorch中如何基于批量张量索引另一张量的最后两个维度?

解决PyTorch中按batch对应索引提取张量元素的问题

你的代码没法得到想要的结果,原因是index_select是对整个维度全局选取索引,但你需要的是每个batch对应自己的专属索引,而不是所有batch都用同一组索引。下面是两种可行的实现方式:

方法一:高级索引(最直观)

直接利用PyTorch的高级索引,匹配每个batch和通道对应的位置,再扩展维度得到目标形状:

import torch

A = torch.randn(2, 3, 55, 45)
B = torch.LongTensor([[2, 5], [10, 20]])

# 生成batch和通道的索引矩阵,用于匹配每个(batch, channel)对
batch_idx = torch.arange(A.size(0)).unsqueeze(1)  # shape [2,1]
channel_idx = torch.arange(A.size(1)).unsqueeze(0)  # shape [1,3]

# 提取对应位置的元素,此时形状为[2,3]
selected = A[batch_idx, channel_idx, B[:,0].unsqueeze(1), B[:,1].unsqueeze(1)]
# 扩展最后两个维度为1,得到目标形状[2,3,1,1]
C = selected.unsqueeze(-1).unsqueeze(-1)

方法二:使用torch.gather(更灵活)

如果需要更通用的维度索引,可以用gather函数,通过扩展索引的维度来匹配原张量的广播规则:

import torch

A = torch.randn(2, 3, 55, 45)
B = torch.LongTensor([[2, 5], [10, 20]])

# 将索引扩展为[2,1,1,1],方便广播到原张量的前两维
idx_dim2 = B[:, 0].unsqueeze(1).unsqueeze(-1).unsqueeze(-1)
idx_dim3 = B[:, 1].unsqueeze(1).unsqueeze(-1).unsqueeze(-1)

# 先在维度2上提取对应索引的元素,得到shape [2,3,1,45]
temp = torch.gather(A, dim=2, index=idx_dim2.expand(-1, 3, 1, 45))
# 再在维度3上提取对应索引的元素,得到目标shape [2,3,1,1]
C = torch.gather(temp, dim=3, index=idx_dim3.expand(-1, 3, 1, 1))

两种方法都能得到你需要的[2,3,1,1]形状的张量C,其中每个C[i,j,0,0]对应A[i,j,B[i,0],B[i,1]]的值。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 02:00:21