当前位置: 首页>>代码示例 >>用法及示例精选 >>正文


Python PyTorch Tensor.index_copy_用法及代码示例


本文简要介绍python语言中 torch.Tensor.index_copy_ 的用法。

用法:

Tensor.index_copy_(dim, index, tensor) → Tensor

参数

  • dim(int) -索引的维度

  • index(LongTensor) - tensor 的索引可供选择

  • tensor(Tensor) -包含要复制的值的张量

通过按 index 中给出的顺序选择索引,将 tensor 的元素复制到 self 张量中。例如,如果 dim == 0index[i] == j ,则将 tensor 的第 i 行复制到 self 的第 j 行。

tensor 的第 dim 维度必须与 index 的长度(必须是向量)具有相同的大小,并且所有其他维度必须匹配 self ,否则将引发错误。

注意

如果 index 包含重复条目,则 tensor 中的多个元素将被复制到 self 的同一索引中。结果是不确定的,因为它取决于最后出现的副本。

例子:

>>> x = torch.zeros(5, 3)
>>> t = torch.tensor([[1, 2, 3], [4, 5, 6], [7, 8, 9]], dtype=torch.float)
>>> index = torch.tensor([0, 4, 2])
>>> x.index_copy_(0, index, t)
tensor([[ 1.,  2.,  3.],
        [ 0.,  0.,  0.],
        [ 7.,  8.,  9.],
        [ 0.,  0.,  0.],
        [ 4.,  5.,  6.]])

相关用法


注:本文由纯净天空筛选整理自pytorch.org大神的英文原创作品 torch.Tensor.index_copy_。非经特殊声明,原始代码版权归原作者所有,本译文未经允许或授权,请勿转载或复制。