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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 16:32:32