PyTorch张量多维选择的高效实现方案问询
高效实现PyTorch张量多维选择
你的问题核心是避免创建超大辅助张量来实现索引,以下是两种高效的解决方案:
方法一:利用高级索引(最优)
直接通过广播式的高级索引定位目标元素,完全不需要额外辅助张量,内存和计算效率最高:
import torch B, V, d = 2, 20000, 64 N, k = 30000, 10 a = torch.rand(B, V, d) b = torch.randint(0, V, (B, N, k)) # 直接通过高级索引获取结果 c = a[torch.arange(B)[:, None, None], b, :]
原理说明:
torch.arange(B)[:, None, None]生成形状为[B, 1, 1]的索引,和b(形状[B, N, k])广播后,形成[B, N, k]的批次索引,对应每个元素所属的batch。- 最终索引组合
[batch_idx, b_idx, :]直接从a中提取每个batch下、b指定的V维度索引对应的所有d维度元素,输出形状为[B, N, k, d],和原方法结果一致。
方法二:用expand替代repeat减少内存占用
如果更习惯用gather操作,可以将原方法中的repeat替换为expand,因为expand不会复制数据,只是逻辑上扩展维度,避免创建超大张量:
# 用expand替代repeat,不复制数据 help_1 = a[:, None, :, :].expand(-1, N, -1, -1) # 形状[B, N, V, d],但内存复用a的数据 help_2 = b[:, :, :, None].expand(-1, -1, -1, d) c = torch.gather(help_1, dim=2, index=help_2)
对比原方法:
原方法中repeat(1, N, 1, 1)会将a在N维度复制N次,生成的help_1内存占用是a的N倍;而expand只是扩展视图,内存占用和a完全相同,大幅降低内存开销。
两种方法都能得到和原实现一致的结果,但方法一的内存效率最高,推荐优先使用。
内容的提问来源于stack exchange,提问作者YuxuanSnow
相关产品推荐
相关产品推荐

