PyTorch中如何用形状[b,k]的张量索引[b,m,n]形状的张量
在PyTorch中实现按batch维度索引张量
你可以通过以下几种方法实现需求,核心是让索引张量的形状适配PyTorch的索引规则:
方法1:使用torch.take_along_dim(推荐,PyTorch 1.10+)
这个函数专门用于沿指定维度选取元素,语义更直观,不需要复杂的形状扩展:
import torch # 构造示例张量 b, m, n, k = 2, 5, 3, 3 A = torch.randn(b, m, n) # shape: [b, m, n] B = torch.randint(0, m, (b, k)) # shape: [b, k],索引值范围0~m-1 # 将索引张量扩展为[b, k, 1],自动广播匹配n维度 indices = B.unsqueeze(-1) result = torch.take_along_dim(A, indices, dim=1) # result形状为[b, k, n],符合要求
方法2:使用torch.gather
需要将索引张量扩展为与输入张量A的形状匹配(仅在索引维度dim=1上长度不同):
# 将B扩展为[b, k, n],每个n维度位置复用相同的索引值 indices = B.unsqueeze(-1).expand(-1, -1, n) result = torch.gather(A, dim=1, index=indices) # result形状为[b, k, n]
方法3:使用高级索引
通过构造batch维度的索引,直接进行多维索引:
# 构造batch索引,形状[b, k],每个batch对应自身的索引 batch_idx = torch.arange(b).unsqueeze(1).expand(-1, k) result = A[batch_idx, B, :] # result形状为[b, k, n]
说明
torch.index_select/torch.take仅支持全局一维索引,无法直接处理batch维度的独立索引,所以不适用。- 之前使用
torch.gather失败是因为未将索引张量扩展到与输入张量的所有非索引维度匹配,扩展后即可正常工作。
内容的提问来源于stack exchange,提问作者cloverse
相关产品推荐
相关产品推荐

