如何在PyTorch中无中间分配或循环复制指定索引张量元素
问题描述
给定以下代码:
import torch a: torch.Tensor b: torch.Tensor assert a.shape[1:] == b.shape[1:] idx = torch.randint(b.shape[0], [a.shape[0]])
需要执行操作 b[...] = a[idx],但不希望产生 a[idx] 带来的中间缓冲区,也不希望对 idx 进行循环遍历,该如何实现?
解决方案
可以直接使用 torch.index_select 并指定 out 参数,将索引后的结果直接写入 b,完全避免中间张量的创建:
torch.index_select(a, 0, idx, out=b)
说明
torch.index_select用于在指定维度上按索引提取元素,这里指定维度0(对应张量的第一维度),使用idx作为索引张量。- 通过
out=b参数,操作会直接将结果写入b的内存空间,不需要额外创建a[idx]这样的中间缓冲区,满足内存高效的需求。 - 该操作是向量化实现,不需要手动循环遍历
idx,性能和原生索引操作一致。
内容的提问来源于stack exchange,提问作者yuri kilochek
相关产品推荐
相关产品推荐

