如何在PyTorch中高效实现批量张量索引(替代循环)
PyTorch批量张量索引的高效实现
基础场景回顾
当源张量source形状为(3, 2),索引张量index形状为(3, 3)时,直接通过source[index]即可得到形状为(3, 3, 2)的结果,示例如下:
import torch source = torch.tensor([[1, 6], [2, 3], [8, 0]]) index = torch.tensor([[2, 1, 2], [1, 1, 2], [2, 0, 0]]) output = source[index] # output形状: (3, 3, 2)
批量场景需求
当批量大小为2时,source形状为(2, 3, 2),index形状为(2, 3, 3),期望得到形状为(2, 3, 3, 2)的结果,无需循环即可高效实现。
高效实现方法
方法1:高级索引(最直观)
通过生成批量维度的索引,与index配合完成索引操作:
# 构造批量示例数据 source = torch.tensor([ [[1, 6], [2, 3], [8, 0]], [[4, 5], [6, 7], [9, 1]] ]) # shape: (2, 3, 2) index = torch.tensor([ [[2, 1, 2], [1, 1, 2], [2, 0, 0]], [[1, 0, 2], [2, 2, 0], [0, 1, 1]] ]) # shape: (2, 3, 3) # 生成批量维度索引,形状与index一致 batch_idx = torch.arange(source.size(0)).unsqueeze(1).unsqueeze(2).expand_as(index) # 执行索引 result = source[batch_idx, index] # result形状: (2, 3, 3, 2)
方法2:使用torch.gather
通过调整张量维度,配合gather函数完成按维度索引:
# 基于上述相同的source和index # 将source扩展为(2, 3, 1, 2),为索引预留维度 source_expanded = source.unsqueeze(2) # 将index扩展为(2, 3, 3, 2),匹配source的最后一维长度 index_expanded = index.unsqueeze(-1).expand(-1, -1, -1, source.size(-1)) # 在dim=1维度上执行gather result = torch.gather(source_expanded, dim=1, index=index_expanded) # result形状: (2, 3, 3, 2)
两种方法均无需循环,完全利用PyTorch的内置张量操作实现,效率远高于循环遍历批量元素。
内容的提问来源于stack exchange,提问作者Pranav Jadhav
相关产品推荐
相关产品推荐

