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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 01:22:44