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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.14 04:11:09